AutoGluon 1.0 版本深度解读:Dynamic Stacking 与 Zeroshot-HPO 驱动的性能跃升
【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon
AutoGluon 1.0 是该项目迈向正式里程碑的关键版本,发布说明(docs/whats_new/v1.0.0.md)详尽记录了它在 Tabular、AutoMM(多模态)、TimeSeries(时间序列)三大模块上的核心增强、完整弃用 API 清单与配套论文。本文以该发布说明为主体,结合当前仓库源码(预设配置、Zeroshot 组合、问题类型常量、指标实现等),帮助读者快速掌握 1.0 的技术变化、关键参数语义与升级迁移路径,可用于模型选型、代码迁移与版本对比研究。
版本总览:四年开发,一次全面换代
AutoGluon 1.0 于 2023 年底发布,官方将其定义为"用 3 行代码实现快速且准确的机器学习"愿景的集中体现。该版本包含 223 个提交,来自 17 位贡献者,并支持 Python 3.8、3.9、3.10 与 3.11。
需要特别注意的是:加载旧版本 AutoGluon 训练出的模型是不受支持的,升级后必须使用 1.0 重新训练模型。这一兼容性边界在发布说明中被明确标注,是迁移前必须评估的关键约束。
依赖升级一览
1.0 对核心依赖做了整体升级,发布说明中列出的版本范围如下:
| 依赖 | 版本范围 | 影响面 |
|---|---|---|
| torch | >=2.0,<2.2 | 深度学习后端全面升到 2.x 时代 |
| numpy | >=1.21,<1.29 | 数值计算基础 |
| pandas | >=2.0,<2.2 | DataFrame 处理(2.x 大版本) |
| scikit-learn | >=1.3,<1.5 | 传统机器学习模型与工具 |
| scipy | >=1.5.4,<1.13 | 科学计算 |
| LightGBM | >=3.3,<4.2 | GBDT 模型 |
| XGBoost | >=1.6,<2.1 | GBDT 模型 |
| Pillow | >=10.0.1,<11 | 图像处理 |
此外,通用模块还新增了系统信息日志工具(system info logging utility),并在 AutoMM 侧将 lightning 升级到 2.0、torchmetrics 升级到 1.0,为多模态训练的稳定性打下基础。
Tabular:两项核心创新带来预测质量跃升
发布说明将 1.0 称为"自 2020 年 3 月原始 AutoGluon 论文以来,表格数据领域最大的一次 SOTA 跃迁",其核心归功于两项特性:Dynamic Stacking(动态堆叠)与Zeroshot-HPO 学习到的超参数组合(portfolio)。二者叠加使得 1.0 相较 0.8 具备 75% 的胜率(win-rate),同时推理更快、磁盘占用更小、稳定性更高。
Dynamic Stacking:缓解堆叠过拟合
传统多层堆叠(multi-layer stacking)中,高层模型使用的训练数据是由低层模型的预测拼接而成的,这天然引入"标签泄漏式的过拟合"——即低层模型对自身训练样本的预测过于乐观,导致高层模型学到的权重失真。Dynamic Stacking 通过在训练阶段动态判断哪些样本可以被安全地用于堆叠,从而缓解这一问题。
从源码看,dynamic_stacking是TabularPredictor.fit的一个显式参数,签名位于 predictor.py,类型为bool | str,默认值为False。其行为关键点(对应_sanitize_stack_args逻辑,见 predictor.py):
- 当传入字符串
"auto"时,AutoGluon 会根据use_bag_holdout等验证方案自动决定是否启用:若use_bag_holdout被禁用则启用 dynamic stacking,否则跳过; - 当
dynamic_stacking=True但num_stack_levels < 1时,会自动强制降级为False(因为不存在可优化的堆叠层); - 启用后(
_dynamic_stacking方法,见 predictor.py),系统会动态调整num_stack_levels与剩余time_limit,若时间不足还会提示设置dynamic_stacking=False; - Dynamic Stacking 仅在首次
fit时检测堆叠过拟合,fit_extra场景不支持(见 predictor.py)。
dynamic_stacking还支持通过ds_*前缀的 keyword arguments 做细粒度控制(见 predictor.py)。
Zeroshot-HPO 组合:用 1 万次实验换来的超参数先验
第二个创新来自 TabRepo 集成仿真库与 Zeroshot-HPO 技术:在大量数据集上预先进行零样本超参数搜索,将表现优异的模型配置整理成一个"组合(portfolio)",作为新数据集的默认超参数起点。仓库中的实现位于 tabular/src/autogluon/tabular/configs/zeroshot/zeroshot_portfolio_2023.py,其文件头注释明确写道:
"Portfolio learned from zeroshot-HPO using TabRepo. Contains the default AutoGluon-Tabular hyperparameters as well as up to 100 learned model configs."
从该文件内容可以看到,组合中为NN_TORCH等模型配备了多组带优先级的完整超参(如activation、dropout_prob、hidden_size、learning_rate、num_layers、weight_decay、use_batchnorm),并通过ag_args.name_suffix与priority控制集成时的排序权重。这正是 1.0 能同时在准确率与稳定性上超越 0.8 的"配方"来源。
预设(Preset)中的落地方式
Zeroshot 组合被接入到best_quality与high_quality预设中,配置定义于 presets_configs.py:
best_quality = { "auto_stack": True, "dynamic_stacking": "auto", "hyperparameters": "zeroshot", "time_limit": 3600, } high_quality = { "auto_stack": True, "dynamic_stacking": "auto", "hyperparameters": "zeroshot", "time_limit": 3600, "refit_full": True, "set_best_to_refit_full": True, "save_bag_folds": False, }其中hyperparameters: "zeroshot"即指向上述组合(导入见 hyperparameter_configs.py)。high_quality与best_quality的差异在于:它通过refit_full+set_best_to_refit_full+save_bag_folds: False在保持高精度的同时将推理速度和磁盘占用优化约 8 倍(发布说明原文表述为"8x faster inference and 8x less disk usage")。good_quality预设则使用hyperparameters: "light"进一步换取训练速度。用户也可以直接在fit中通过hyperparameters指定其他字符串(如"default"、"light")或自定义字典。
其他 Tabular 新特性与修复
- 实验性 scikit-learn API:可通过
from autogluon.tabular.experimental import TabularClassifier, TabularRegressor使用与 sklearn 兼容的分类器/回归器封装。仓库实现见 _tabular_classifier.py 与 _tabular_regressor.py,二者分别继承BaseEstimator、ClassifierMixin/RegressorMixin与ScikitMixin; - 新增
predictor.model_failures()与predictor.simulation_artifact()(后者用于与 TabRepo 集成); - 新增增强版 FT-Transformer(在 AutoMM 的 tabular backbone 中替换了 MLP);
- 性能改进:FastAI 回归输出裁剪、Skip-connection Weighted Ensemble、用 Ray 子进程做顺序拟合以修复内存泄漏(发布说明称其在 AutoML Benchmark 上实现0 任务失败)、低内存场景的动态并行折叠(dynamic parallel folds)支持;
- 稳定性:多层堆叠现在产生确定性结果;修复了模型内存占用计算错误、bagging 时
infer_limit使用错误、FastAI 罕见崩溃等问题。
推理速度与精度的取舍:infer_limit
发布说明强调 1.0 是"在保持 SOTA 精度的前提下推理吞吐最快"的 AutoML 系统:用户可通过fit的infer_limit参数在精度与推理速度之间做显式权衡(发布说明也记录了 bagging 场景下infer_limit使用错误的修复)。
OpenML AutoML Benchmark 结果回顾
发布说明引用了 OpenML 于 2023 年 11 月 16 日发布的官方 2023 AutoML Benchmark 结果(1040 个任务)。根据其表述,AutoGluon 1.0 对传统表格模型胜率在 95% 以上,其中对 LightGBM 胜率 99%、对 XGBoost 胜率 100%;对其他 AutoML 系统胜率在 82%~94% 之间;且在 63% 的任务中取得第一名(第二名 lightautoml 为 12%,AutoGluon 0.8 此前为 48%)。完整对比表如下(数据源自发布说明):
| 方法 | AG 胜率 | AG 损失改进 | Rescaled Loss | 平均排名 | 冠军占比 |
|---|---|---|---|---|---|
| AutoGluon 1.0 (Best, 4h8c) | - | - | 0.04 | 1.95 | 63% |
| lightautoml (2023, 4h8c) | 84% | 12.0% | 0.2 | 4.78 | 12% |
| H2OAutoML (2023, 4h8c) | 94% | 10.8% | 0.17 | 4.98 | 1% |
| FLAML (2023, 4h8c) | 86% | 16.7% | 0.23 | 5.29 | 5% |
| MLJAR (2023, 4h8c) | 82% | 23.0% | 0.33 | 5.53 | 6% |
| autosklearn (2023, 4h8c) | 91% | 12.5% | 0.22 | 6.07 | 4% |
| GAMA (2023, 4h8c) | 86% | 15.4% | 0.28 | 6.13 | 5% |
| CatBoost (2023, 4h8c) | 95% | 18.2% | 0.28 | 6.89 | 3% |
| TPOT (2023, 4h8c) | 91% | 23.1% | 0.4 | 8.15 | 1% |
| LightGBM (2023, 4h8c) | 99% | 23.6% | 0.4 | 8.95 | 0% |
| XGBoost (2023, 4h8c) | 100% | 24.1% | 0.43 | 9.5 | 0% |
| RandomForest (2023, 4h8c) | 97% | 25.1% | 0.53 | 9.78 | 1% |
需要说明的是:上表为发布说明中 AutoGluon 团队基于 OpenML 基准的自我报告数据,读者可将其作为版本性能定位的参考。此外,发布说明指出 AutoGluon 1.0 相较 0.8 还有平均 7.4% 的损失改进,且凭借低内存训练中新的 Ray 子进程方案实现了 0 任务失败。
AutoMM(AutoGluon Multimodal):基础模型微调的全面扩展
AutoMM 的目标是"用三行代码微调基础模型(foundation models)",与 HuggingFace Transformers、TIMM、MMDetection 等模型库无缝集成,支持图像、文本、表格、文档数据及其任意组合。与主要聚焦表格分类/回归的其他开源 AutoML 工具(如 AutoSklearn、LightAutoML、H2OAutoML、FLAML、MLJAR、TPOT、GAMA)相比,AutoMM 是当时唯一同时覆盖多模态数据、多样化任务与基础模型体系(含深度学习模型)的 AutoML 系统。发布说明给出的能力对比矩阵如下:
| 系统 | 图像 | 文本 | 表格 | 文档 | 任意组合 | 分类 | 回归 | 目标检测 | 语义匹配 | 命名实体识别 | 图像分割 | 传统模型 | 深度学习模型 | 基础模型 |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| LightAutoML | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | |||||||
| H2OAutoML | ✓ | ✓ | ✓ | ✓ | ||||||||||
| FLAML | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | |||||||
| MLJAR | ✓ | ✓ | ✓ | ✓ | ||||||||||
| AutoSklearn | ✓ | ✓ | ✓ | ✓ | ✓ | |||||||||
| GAMA | ✓ | ✓ | ✓ | ✓ | ||||||||||
| TPOT | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ||||||||
| AutoMM | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ |
新问题类型:语义分割(Semantic Segmentation)
1.0 引入全新问题类型semantic_segmentation,用于对 Segment Anything Model (SAM) 进行三行代码微调。仓库中该问题类型定义于 constants.py,并配有独立的 learner 实现 semantic_segmentation.py,默认评估指标为 IOU/BER/SM(见 constants.py)。发布说明要点:
- 覆盖自然图像、农业、遥感、医疗等多样化领域基准;
- 使用参数高效微调(PEFT)中的 LoRA 方法,在基准测试中一致优于 VPT、adaptor、BitFit、SAM-adaptor、LST 等替代方案;
- 默认使用 SAM-ViT Huge 骨干,要求 GPU 显存大于 25GB;
- 配套教程见 docs/tutorials/multimodal/image_segmentation/beginner_semantic_seg.ipynb,完整示例见 examples/automm/Conv-LoRA/run_semantic_segmentation.py。
新问题类型:少样本分类(Few Shot Classification)
新增few_shot_classification问题类型(常量见 constants.py),利用图像/文本基础模型提取特征并训练 SVM 分类器,适用于小样本学习场景。配套教程见 docs/tutorials/multimodal/advanced_topics/few_shot_learning.ipynb。同时,旧的FewShotSVMPredictor被弃用。
性能与工程改进
- torch.compile 支持:用于加速训练(实验性,需 torch >= 2.2);
- 默认骨干升级:图像默认骨干升级后在图像基准上实现 100% 胜率;表格默认骨干由 MLP 替换为 FT-Transformer,在文本+表格基准上实现 67% 胜率;二者叠加在文本+表格+图像基准上实现 62% 胜率;
- 稳定性:开启严格的多 GPU CI 测试并修复多项多 GPU 问题;新增 DDP 策略的多 GPU 推理支持;
- 易用性:支持自定义评估指标(将 metric 对象传入
eval_metric);笔记本中支持多 GPU 训练(实验性);新增hf_text.use_fast选项控制快速 tokenizer 使用;补充f1_macro、f1_micro、f1_weighted回退评估指标; - 可扩展性:引入新的 learner 类设计(learners/base.py),便于扩展新任务与新模态;FT-Transformer、目标检测/语义分割/NER 可视化器均完成重构。
TimeSeries:鲁棒性、新模型与新指标
发布说明指出,1.0 的 TimeSeries 模块在易用性、性能与鲁棒性上均有大量改进,并宣称在预测准确率上对主流预测框架取得 70%+ 胜率(该表述对应 AutoML Conference 2023 论文)。
数据鲁棒性
TimeSeriesPredictor现在可以处理所有 pandas 频率的数据、不规则时间戳以及用NaN表示的缺失值。TimeSeriesDataFrame也支持在from_path/from_data_frame构造时直接传入静态特征。
新模型
- 间歇性需求预测模型(基于 conformal prediction):
ADIDA、CrostonClassic、CrostonOptimized、CrostonSBA、IMAPA; - 来自 GluonTS 的
WaveNet与NPTS; - 新基线模型:
Average、SeasonalAverage、Zero。
同时,DirectTabular改为基于mlforecast后端实现(与RecursiveTabular一致),RecursiveTabular与DirectTabular的训练/预测速度更快、内存占用更低。
新指标与自定义指标
新增点预测指标WAPE、RMSSE与SQL。这些指标已在仓库中注册实现(导入与别名见 timeseries/metrics/init.py,如weighted_absolute_percentage_error映射到WAPE)。同时,1.0 支持自定义预测指标,并允许向TimeSeriesPredictor.evaluate传入多个评估指标。注意:旧指标名"mean_wQuantileLoss"已更名为"WQL"。
高级交叉验证选项
TimeSeriesPredictor.fit新增两个参数:
refit_every_n_windows:避免为每个验证窗口重复训练模型;val_step_size:调整验证窗口之间的步长。
多窗口/最后窗口 splitter 类(MultiWindowSplitter、LastWindowSplitter)被弃用,改为上述参数或ExpandingWindowSplitter。
其他增强
- 启用 Ray Tune 支持深度学习预测模型超参调优;
- 低时间预算下通过新预设与训练时间分配逻辑实现更准确的预测;
- GluonTS 模型支持早停并提升推理速度;通过将 import 移入模型类内部减少
autogluon.timeseries的导入时间; TimeSeriesPredictor的 API 与TabularPredictor对齐,移除弃用方法。
EDA 模块说明
EDA 模块因仍需更多开发工作,不在 1.0 中发布。如有需要请继续使用"autogluon.eda==0.8.2",官方将在其就绪后另行公告。
弃用(Deprecations)与升级迁移清单
1.0 是一次大规模 API 整理,发布说明给出了完整的迁移映射,升级时务必对照修改代码。
通用模块
autogluon.core.spaces已弃用,请改用autogluon.common.spaces。
TabularPredictor 方法重命名
以下旧方法仍可用但会输出警告日志,计划在AutoGluon 1.2 移除:
| 旧方法(已弃用) | 新方法 |
|---|---|
predictor.get_model_names() | predictor.model_names() |
predictor.get_model_names_persisted() | predictor.model_names(persisted=True) |
predictor.compile_models() | predictor.compile() |
predictor.persist_models() | predictor.persist() |
predictor.unpersist_models() | predictor.unpersist() |
predictor.get_model_best() | predictor.model_best |
predictor.get_pred_from_proba() | predictor.predict_from_proba() |
predictor.get_oof_pred_proba() | predictor.predict_proba_oof() |
predictor.get_oof_pred() | predictor.predict_oof() |
predictor.get_model_full_dict() | predictor.model_refit_map() |
predictor.get_size_disk() | predictor.disk_usage() |
predictor.get_size_disk_per_file() | predictor.disk_usage_per_file() |
leaderboard()/evaluate()/evaluate_predictions()的silent参数 | 改用display,默认值为False |
AutoMM 弃用
FewShotSVMPredictor弃用,改用新的few_shot_classification问题类型;AutoMMPredictor弃用,改用MultiModalPredictor;MultiModalPredictor.fit的config参数弃用;MultiModalPredictor初始化 API 的init_scratch与pipeline参数弃用。
TimeSeries 弃用
TimeSeriesPredictor(ignore_time_index: bool)参数弃用:若数据包含不规则时间戳,应通过data.convert_frequency(freq)转换为规则频率,或在创建 predictor 时指定TimeSeriesPredictor(freq=freq);predictor.evaluate()现在返回字典(此前返回浮点数);predictor.score()→predictor.evaluate();get_model_names()→model_names();get_model_best()→model_best;- 指标
"mean_wQuantileLoss"更名为"WQL"; leaderboard()的silent参数改用display(默认False);fit中hyperparameters传字符串时,仅支持"default"、"light"、"very_light";TimeSeriesDataFrame.to_regular_index()→convert_frequency();get_reindexed_view()弃用;- 基于 MXNet 的模型全部移除(
DeepARMXNet、MQCNNMXNet、MQRNNMXNet、SimpleFeedForwardMXNet、TemporalFusionTransformerMXNet、TransformerMXNet); - 基于 Statsmodels 的统计模型(
ARIMA、Theta、ETS)替换为 StatsForecast 版本,超参数名称发生变化; DirectTabular改用 mlforecast 后端,大部分超参数名称变化;TimeSeriesEvaluator弃用,改用autogluon.timeseries.metrics中的指标;MultiWindowSplitter与LastWindowSplitter弃用,改用num_val_windows、val_step_size参数或ExpandingWindowSplitter。
配套论文
发布说明汇总了 1.0 相关的五篇论文,可作为深入了解算法原理的入口:
- AutoGluon-TimeSeries: AutoML for Probabilistic Time Series Forecasting(AutoML Conference 2023):对 DeepAR、TFT、AutoARIMA、AutoETS、AutoPyTorch 等框架的基准显示,AutoGluon 在点预测与概率预测上均达 SOTA,且对"事后最优模型组合"取得 65% 胜率;
- TabRepo: A Large Scale Repository of Tabular Model Evaluations and its AutoML Applications(arXiv 2311.02971):表格 Zeroshot-HPO 集成仿真库,是 1.0 性能提升的关键支撑;
- XTab: Cross-table Pretraining for Tabular Transformers(ICML 2023):表格 Transformer 预训练,可匹配 XGBoost/LightGBM 性能,尚未集成进 AutoGluon,计划未来版本引入;
- Learning Multimodal Data Augmentation in Feature Space(ICLR 2023):特征空间中的多模态数据增强模块(LeMDA),未集成,计划未来引入;
- Data Augmentation for Object Detection via Controllable Diffusion Models(WACV 2024):基于可控扩散模型与 CLIP 的目标检测数据增强流水线,未集成;
- Adapting Image Foundation Models for Video Understanding(ICLR 2023):通过空间/时间/联合适配让冻结的图像基础模型获得时空推理能力。
结语与升级建议
AutoGluon 1.0 通过 Dynamic Stacking 与 Zeroshot-HPO 组合两项创新将表格学习的精度与稳定性推上新的台阶,同时让 AutoMM 覆盖语义分割与少样本分类两大新任务,并让 TimeSeries 模块在数据鲁棒性、模型库与交叉验证能力上全面升级。对于正在使用旧版本的用户,建议按以下路径升级:
- 检查依赖环境(Python 3.8–3.11,torch >= 2.0)与上述依赖版本范围;
- 对照弃用清单批量替换已弃用 API;
- 注意 TimeSeries 模型(MXNet 系、Statsmodels 系、DirectTabular 后端)的移除与超参数变化;
- 旧模型文件无法在 1.0 加载,需重新训练;
- 在
fit中优先使用best_quality/high_quality预设(其内部已自动启用dynamic_stacking: "auto"与hyperparameters: "zeroshot"),并按需用infer_limit权衡推理速度。
【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考