AutoGluon:多模态自动机器学习库
bash pip install autogluon
### 完整安装(支持图像、文本和多模态)
```bash
pip install autogluon[all]
对于图像任务,还需额外安装支持 GPU 的 MXNet 或 PyTorch(根据你的硬件),一般通过 pip install mxnet-cu112 或 pip install torch torchvision 完成。
验证安装
import autogluon as ag
print(ag.__version__)
表格预测:5 分钟上手 AutoGluon
表格数据(结构化数据)是最常见的应用场景。以下示例使用内置的 Titanic 数据集演示分类任务。
加载数据
from autogluon.tabular import TabularDataset, TabularPredictor
train_data = TabularDataset('https://autogluon.s3.amazonaws.com/datasets/Inc/train.csv')
test_data = TabularDataset('https://autogluon.s3.amazonaws.com/datasets/Inc/test.csv')
label = 'class' # 目标列名
训练模型
predictor = TabularPredictor(label=label).fit(train_data)
训练过程中,AutoGluon 会自动识别问题类型(二分类、多分类或回归),进行缺失值填充、特征工程,并依次训练多个基础模型(如 LightGBM、CatBoost、XGBoost、神经网络等),最后将它们集成。
评估与预测
# 在测试集上评估
performance = predictor.evaluate(test_data)
# 对新样本预测
y_pred = predictor.predict(test_data.drop(columns=[label]))
evaluate() 会输出准确率、对数损失等指标(分类)或 RMSE(回归)。predict() 返回预测类别或数值。
保存与加载模型
predictor.save('my_model') # 保存至文件夹
loaded_predictor = TabularPredictor.load('my_model')
内置评估与排行榜
训练结束后,predictor.leaderboard() 会展示所有模型的性能排序,帮助你快速了解哪个模型表现更好。
leaderboard = predictor.leaderboard(test_data, silent=True)
print(leaderboard)
你也可以通过 predictor.feature_importance() 查看特征重要性,辅助理解数据。
图像分类:一行代码获得高精度模型
AutoGluon 封装了多种先进的图像分类模型,并自动进行数据增强、学习率调度。
准备图像数据集
图像数据集应按 train/类别名/*.jpg 和 test/类别名/*.jpg 的组织方式存放,或直接使用 ImageFolder 格式。
from autogluon.vision import ImagePredictor, ImageDataset
train_dataset = ImageDataset.from_folders('path/to/train/')
test_dataset = ImageDataset.from_folders('path/to/test/')
训练与预测
predictor = ImagePredictor()
predictor.fit(train_dataset, hyperparameters={'epochs': 10})
# 评估准确率
test_acc = predictor.evaluate(test_dataset)
print(f'测试准确率: {test_acc}')
# 单张图片预测
pred = predictor.predict('path/to/image.jpg')
你可以传入自定义模型(如 resnet50、efficientnet_b0)或让 AutoGluon 自动搜索最优架构。
文本分类:自然语言任务一键搞定
AutoGluon 支持对原始文本进行分类、情感分析等任务,内部使用预训练的 Transformer 模型(如 ELECTRA、DeBERTa)并进行微调。
数据格式
数据通常是一个 CSV 文件,包含文本列和标签列。
from autogluon.text import TextPredictor
train_data = pd.read_csv('train.csv') # 包含 'text' 和 'label' 两列
test_data = pd.read_csv('test.csv')
训练流程
predictor = TextPredictor(label='label', eval_metric='acc')
predictor.fit(train_data, hyperparameters='multilingual', # 英文可用 'default'
time_limit=600) # 时间限制(秒)
hyperparameters 可选择 'default'(英文高性能)、'multilingual'(多语言支持)或自定义模型列表。
预测
predictions = predictor.predict(test_data['text'])
score = predictor.evaluate(test_data)
多模态融合:同时利用图像与文本
当数据中包含图像和文本字段时,AutoGluon 的多模态模块可自动将这些信息融合到一个模型中。
数据结构示例
假设你的 DataFrame 包含:
image_path:图像文件路径description:文本描述label:分类标签
from autogluon.multimodal import MultiModalPredictor
train_data = pd.read_csv('multimodal_train.csv')
predictor = MultiModalPredictor(label='label')
predictor.fit(train_data, time_limit=600)
训练过程中,AutoGluon 会自动对图像用 CNN 提取特征,对文本用 Transformer,并通过可学习的融合层将它们结合,从而获得比单模态更强的性能。
预测
test_data = pd.read_csv('multimodal_test.csv')
predictions = predictor.predict(test_data)
模型可解释性:看清模型如何决策
AutoGluon 提供了多种解释工具,特别适合需要模型透明度的场景。
表格模型解释
# 全局特征重要性
fi = predictor.feature_importance(test_data)
print(fi)
# 局部解释(SHAP 值,需安装 shap 库)
explanations = predictor.explain_rows(test_data.head(), method='shap')
图像模型解释
from autogluon.vision import ImagePredictor
predictor = ImagePredictor.load('my_image_model')
saliency = predictor.explain('image.jpg')
saliency.plot() # 绘制显著图
自定义配置与高级控制
虽然 AutoGluon 的默认设置能应对多数场景,但你也可以精细地调节各个环节。
表格预测常用参数
predictor = TabularPredictor(label=label, eval_metric='roc_auc').fit(
train_data,
presets='best_quality', # 预设:'medium_quality'、'high_quality'、'best_quality'
hyperparameters={
'GBM': {'num_boost_round': 500},
'NN_TORCH': {},
'RF': {},
},
num_bag_folds=5, # 集成折数(提升稳定性)
time_limit=3600,
verbosity=2 # 日志详细程度
)
自定义搜索空间(仅文本示例)
from autogluon.text import TextPredictor
predictor = TextPredictor(label='label')
predictor.fit(train_data,
hyperparameters={
'models': {
'DeBERTa': {'search_space': {'learning_rate': 2e-5}}
}
})
超参数调优
你可以使用 hyperparameter_tune_kwargs 参数配合 'auto' 策略让 AutoGluon 自动在训练中进一步调参:
predictor.fit(..., hyperparameter_tune_kwargs={'num_trials': 10})
模型部署与导出
训练好的模型可以通过 predictor.save() 保存为文件夹,并可轻松部署到生产环境。
使用 ONNX 导出(表格模型)
predictor.export_onnx(dir_path='onnx_model')