news 2026/9/16 7:00:39

AutoGluon 文本预测实战指南:用 TextPredictor 与多模态配置赢得 NLP 比赛与 GLUE 基准

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
AutoGluon 文本预测实战指南:用 TextPredictor 与多模态配置赢得 NLP 比赛与 GLUE 基准

AutoGluon 文本预测实战指南:用 TextPredictor 与多模态配置赢得 NLP 比赛与 GLUE 基准

【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon

本文围绕 AutoGluon 仓库中 examples/automm/text_prediction 目录下的完整文本预测示例展开,系统讲解如何使用MultiModalPredictor(TextPredictor 是其文本场景的典型用法)以及 AutoGluon Tabular 中的AG_TEXT_NN多模态配置,依次攻克 MachineHack 商品情感分类、图书价格预测、数据科学家薪资预测、Kaggle Mercari 价格建议四类真实比赛任务,并完整复现 GLUE 自然语言理解基准的评测流程。读完本文,你将掌握一套可直接复制运行的比赛级文本建模工作流,理解single / weighted / stacking三种建模模式的区别与底层实现,并能将同一套脚本迁移到自己的文本数据集上。

一、文本预测能力总览:两种 Predictor、一条统一入口

文本预测(Text Prediction)在 AutoGluon 中有两条互补的技术路线,而examples/automm/text_prediction目录下的脚本恰好把两条路线封装成了同一个入口:

  • 单模型路线:直接使用MultiModalPredictor,内部自动完成预训练语言模型的微调(finetune),开箱即用,无需手动指定模型结构。在 run_competition.py 的single模式下可以看到,脚本注释明确写道:"When no embedding is used, we will just use MultiModalPredictor that will train a single model internally."
  • 集成路线:使用TabularPredictor,配合get_hyperparameter_config('multimodal')返回的hyperparameters字典,将文本神经网络(AG_TEXT_NN)与 AutoGluon-Tabular 的经典表格模型(LightGBM、CatBoost 等)组合起来,通过加权集成(weighted)或 5 折 Bagging + 单层 Stacking 提升精度。

两条路线的核心配置文件都来自 tabular/src/autogluon/tabular/configs/hyperparameter_configs.py 中定义的get_hyperparameter_config(config_name)工厂函数,传入'multimodal'即可取回包含AG_TEXT_NN的模型集合。当指定--preset且模式为stacking/weighted时,脚本会把 preset 写入hyperparameters['AG_TEXT_NN']['presets'](见 run_competition.py),从而把多模态模型的训练质量预设透传给内部的文本神经网络。

质量预设(preset)的取值与定义可以在 multimodal/src/autogluon/multimodal/constants.py 中找到:high_qualitymedium_qualitybest_quality,其中medium_quality_faster_train是比赛脚本额外提供的更快训练版本。不同 preset 的差异(如是否启用更大的骨干模型、更长的训练轮数、更强的数据增强)集中实现在 multimodal/src/autogluon/multimodal/utils/presets.py 中,读者可根据算力在"更快训练"与"更高质量"之间权衡。

二、统一比赛脚本 run_competition.py:参数与三种建模模式

目录下的 run_competition.py 是四个 MachineHack / Kaggle 案例共用的驱动脚本,通过--task参数分发到不同的数据加载逻辑,并统一执行"训练 → 预测 → 写提交文件"三步流程。其命令行参数如下(对应源码 get_parser):

参数类型 / 取值默认值说明
--train_filestrNone训练数据文件(CSV / XLSX)
--test_filestrNone测试数据文件
--sample_submissionstrNone比赛提供的提交样例文件
--task必选:product_sentiment/mercari_price/price_of_books/data_scientist_salary-决定数据加载与提交格式
--eval_metricstrNone评测指标,如log_lossr2acc
--modesingle/weighted/stackingsingle建模方式,见下文
--presetmedium_quality_faster_train/high_quality/best_qualityNone多模态模型质量预设
--seedint123随机种子,保证可复现
--exp_dirstrNone模型与提交文件的输出目录

三种模式的核心区别(源码 run):

  1. single:直接MultiModalPredictor.fit(...),内部训练单个预训练语言模型,开销最小、最易上手。
  2. weightedTabularPredictor.fit(...),使用AG_TEXT_NN加表格模型的加权集成,模型之间不做层叠。
  3. stackingTabularPredictor.fit(..., num_bag_folds=5, num_stack_levels=1),即 5 折 Bagging + 1 层 Stacking,把文本模型与表格模型通过一层堆叠集成起来,是比赛脚本默认推荐的高精度方案(README 中多数命令都使用--mode stacking)。

训练完成后,脚本按任务类型生成提交文件:分类任务(如product_sentiment)写出submission.csv(概率形式),回归任务(如mercari_priceprice_of_booksdata_scientist_salary)则先读入--sample_submission再回填预测值。注意:回归任务的预测值在内部经过了log/log10变换,写出前会做对应的指数逆变换(见后文各节),因此最终提交文件中的数值与原始数据同尺度。

三、案例一:MachineHack 商品情感分类(Product Sentiment Classification)

该案例的目标是在 MachineHack 的 "Product Sentiment Classification" 周末黑客松中取得好成绩。比赛数据由Product_Description(商品描述文本)与Product_Type(商品类别)两个特征列和Sentiment标签列构成。README 给出了完整的数据准备与训练命令:

mkdir -p machine_hack_product_sentiment wget https://automl-mm-bench.s3.amazonaws.com/machine_hack_product_sentiment/all_train.csv -O machine_hack_product_sentiment/all_train.csv wget https://automl-mm-bench.s3.amazonaws.com/machine_hack_product_sentiment/test.csv -O machine_hack_product_sentiment/test.csv mkdir -p ag_product_sentiment python3 run_competition.py --train_file machine_hack_product_sentiment/all_train.csv \ --test_file machine_hack_product_sentiment/test.csv \ --task product_sentiment \ --eval_metric log_loss \ --exp_dir ag_product_sentiment \ --mode stacking 2>&1 | tee -a ag_product_sentiment/log.txt

命令执行结束后会在ag_product_sentiment目录生成submission.csv,可直接上传到比赛排行榜。

数据加载逻辑见 load_machine_hack_product_sentiment:训练集只保留Product_DescriptionProduct_Type两个特征列与Sentiment标签列,测试集去掉标签列;预测阶段使用predict_proba(..., as_multiclass=True)输出每个类别的概率,这正是--eval_metric log_loss所要求的提交格式。这里也体现了一个通用技巧:当评测指标是对数损失时,提交概率而非硬标签,由 run_competition.py 中test_probabilities.to_csv(...)完成。

四、案例二:MachineHack 图书价格预测(Predict Price of Book)

第二个案例是 "Predict The Price Of Books" 黑客松的回归任务,目标是根据书名、作者、评论、评分等字段预测图书价格。由于比赛数据以.xlsx格式提供,README 明确要求先安装openpyxl

bash prepare_price_of_books.sh python3 -m pip install openpyxl mkdir -p ag_price_of_books python3 run_competition.py --train_file price_of_books/Participants_Data/Data_Train.xlsx \ --test_file price_of_books/Participants_Data/Data_Test.xlsx \ --sample_submission price_of_books/Participants_Data/Sample_Submission.xlsx \ --task price_of_books \ --eval_metric r2 \ --exp_dir ag_price_of_books \ --mode stacking 2>&1 | tee -a ag_price_of_books/log.txt

prepare_price_of_books.sh会从 AutoGluon 的公开评测数据桶(automl-mm-bench)下载Data.zip并解压到price_of_books目录。脚本结束后,ag_price_of_books目录中会出现submission.xlsx文件。

这个案例在数据预处理上有两个值得借鉴的工程点(见 load_price_of_books):

  1. 文本字段结构化清洗Reviews列原始值是形如"123 out of 5 stars"的字符串,通过ele[:-len(' out of 5 stars')]截断后转为数值;Ratings列原始值是"1,234 customer reviews",先去掉千分位逗号再截断' customer reviews'后缀,同样转为数值。
  2. 标签对数化:对价格做np.log10(Price + 1)变换,把右偏的长尾价格分布拉近正态,显著降低回归误差;写出提交文件时再通过np.power(10, predictions) - 1还原为真实价格(见 run_competition.py)。

README 特别提醒:建议在 p3.2x(或同类 GPU 实例)上运行本实验,因为价格回归需要较充分的模型训练才能达到理想精度。

五、案例三:MachineHack 数据科学家薪资预测(Data Scientist Salary Prediction)

第三个案例是 "Predict The Data Scientists Salary In India Hackathon",根据候选人的工作经验、技能关键词、公司信息等预测其薪资(salary列)。有趣的是该比赛虽然本质是回归问题,README 中却使用--eval_metric acc,从源码 load_data_scientist_salary 看,脚本还会剔除company_name_encoded列,避免模型直接过拟合到公司编码上。

bash prepare_data_scientist_salary.sh python3 -m pip install openpyxl mkdir -p ag_data_scientist_salary python3 run_competition.py --train_file data_scientist_salary/Data/Final_Train_Dataset.csv \ --test_file data_scientist_salary/Data/Final_Test_Dataset.csv \ --sample_submission data_scientist_salary/Data/sample_submission.xlsx \ --task data_scientist_salary \ --eval_metric acc \ --exp_dir ag_data_scientist_salary \ --mode stacking 2>&1 | tee -a ag_data_scientist_salary/log.txt

与图书价格案例相同,prepare_data_scientist_salary.sh负责下载并解压数据,命令完成后ag_data_scientist_salary目录中生成submission.xlsx。由于薪资预测同样对模型容量有要求,README 同样建议在 p3.2x 实例上运行。提交时脚本读取sample_submission.xlsx,将predictor.predict(...)的预测结果回填到salary列后写出(见 run_competition.py)。

六、案例四:Kaggle Mercari 价格建议(Mercari Price Suggestion)

第四个案例来自 Kaggle 的 Mercari Price Suggestion Challenge(二手商品定价),目标是仅凭商品标题、品类、品牌等文本信息预测二手商品价格,README 称之为可达到 Top-5 级别的方案。运行前需要先配置 Kaggle API 以通过命令行下载数据集:

sudo apt install -y p7zip-full bash prepare_mercari_kaggle.sh

prepare_mercari_kaggle.sh依次执行:用 Kaggle API 下载比赛压缩包 → 解压 → 用7za解压train.tsv.7z→ 解压test_stg2.tsv.zipsample_submission_stg2.csv.zip

单模型模式

mkdir -p ag_mercari_price_single python3 run_competition.py --train_file mercari_price/train.tsv \ --test_file mercari_price/test_stg2.tsv \ --sample_submission mercari_price/sample_submission_stg2.csv \ --task mercari_price \ --eval_metric r2 \ --exp_dir ag_mercari_price_single \ --mode single 2>&1 | tee -a ag_mercari_price_single/log.txt

加权集成模式

mkdir -p ag_mercari_price_weighted python3 run_competition.py --train_file mercari_price/train.tsv \ --test_file mercari_price/test_stg2.tsv \ --sample_submission mercari_price/sample_submission_stg2.csv \ --task mercari_price \ --eval_metric r2 \ --exp_dir ag_mercari_price_weighted \ --mode weighted 2>&1 | tee -a ag_mercari_price_weighted/log.txt

Stacking 堆叠模式

mkdir -p ag_mercari_price_stacking python3 run_competition.py --train_file mercari_price/train.tsv \ --test_file mercari_price/test_stg2.tsv \ --sample_submission mercari_price/sample_submission_stg2.csv \ --task mercari_price \ --eval_metric r2 \ --exp_dir ag_mercari_price_stacking \ --mode stacking 2>&1 | tee -a ag_mercari_price_stacking/log.txt

该案例的数据工程最为复杂(见 load_mercari_price_prediction),其要点包括:

  • 层级品类拆分:原始category_name形如"Sports & Outdoors/Outdoor Recreation/Camping & Hiking",脚本用split('/', 2)拆出cat1 / cat2 / cat3三个层级特征,并处理缺失值(None),让模型能同时利用粗细两种粒度的品类信息。
  • 标签对数化:对价格做np.log(price + 1)变换,提交前用np.exp(predictions) - 1还原(见 run_competition.py)。
  • 忽略无信息列train_id被显式排除在特征之外,避免 ID 泄漏。

读者可以对照三种模式各自的效果,在真实数据上验证"集成 > 单模型"的普遍规律,并依据训练耗时选择适合自己算力的模式。

七、GLUE 基准评测:从数据准备到结果复现

README 的最后一节展示了如何用 AutoGluon 文本预测能力解决 GLUE 基准中的全部任务,包括 CoLA、SST-2、MRPC、STS-B、QQP、MNLI(matched / mismatched)、QNLI、RTE、WNLI 等,这节内容在 run_text_prediction.py 与配套脚本中得到了完整实现。

第一步:下载并预处理数据

python3 prepare_glue.py --benchmark glue

prepare_glue.py 是一个功能完整的基准数据管线(部分借鉴自 NLP 社区的 jiant 项目):它定义了GLUE_TASKS/SUPERGLUE_TASKS的任务清单、每个任务专属的读取器(如read_colaread_mrpcread_mnli),通过GLUE_TASK2PATHurl_checksums/glue.txt(SHA-1 校验文件,见 url_checksums)校验下载完整性,并把 TSV 原始数据统一转换为模型友好的parquet 格式。数据默认输出到当前目录的glue/文件夹下。

第二步:跑单模型或 5 折 Stacking 基线

README 说明可以二选一:用单个TextPredictor模型,或使用 AutoGluon Tabular 中的multimodal配置——后者会把TextPredictor与 AutoGluon-Tabular 的表格模型通过单层 Stacking + 5 折 Bagging组合起来。

# Run single model bash run_glue.sh single # Run 5-fold stacking bash run_glue.sh stacking

run_glue.sh 对cola sst mrpc sts qqp qnli rte wnli八个任务逐一调用run_text_prediction.py --do_train,并对 MNLI 分别用 matched / mismatched 的验证与测试集跑mnli_mmnli_mm两个子任务。脚本中每个任务均显式传入--train_file / --dev_file / --test_file,并通过--mode ${MODE}透传 single / stacking 模式。

run_text_prediction.py 是 GLUE 评测的核心执行体,值得关注的设计点:

  • 任务元信息表(源码 TASKS):用字典集中定义每个任务的特征列、标签列、主评测指标与附加指标。例如mrpc使用sentence1 / sentence2双句特征 +label标签 +acc指标,sts使用sentence1 / sentence2+score连续标签 +rmse主指标(附加pearsonrspearmanr)。
  • MRPC / STS 的数据增强(源码 train):README 明确指出"For MRPC and STS, we have manually augmented the training and validation data by shuffling the order of two sentences."。实现上,脚本将两个句子的顺序互换构造一份镜像样本与原始数据拼接,从而让模型对句子顺序不敏感(这两个任务本身对顺序不敏感,但不同数据集存在顺序偏差)。
  • 结果落盘:训练结束后输出dev_prediction.csvtest_prediction.csv,并把验证集指标写入final_model_scores.json,方便批量对比。

第三步:生成 GLUE 提交文件

python3 generate_submission.py --prefix autogluon_text --save_dir submission

generate_submission.py 读取每个任务运行后产生的{prefix}_{task}/test_prediction.csv,拼装成 GLUE 官方提交格式的 TSV 文件(index+ 标签列)。其内部逻辑还包含两个细节:

  • STS-B 特殊处理:预测的相似度分数会被np.clip(predictions, 0, 5)限制在 [0, 5] 合法区间(见 generate_submission.py)。
  • AX(诊断集)处理:AX 没有独立训练数据,脚本直接加载mnli_m的模型检查点做迁移推理(见 generate_submission.py),体现了预训练语言模型在相似任务间迁移的便捷性。

参考结果

README 给出了单模型(Text Single)在各任务验证集上的指标,其中带(*)的 MRPC 与 STS 结果来自句子顺序增强后的数据:

CoLASSTMRPCSTSQQPMNLI-mMNLI-mmQNLIRTEWNLI
指标mccaccaccspearmanrf1accaccaccaccacc
Text (Single) - Validation (*)0.67820.95070.8725 (*)0.9047 (*)0.88660.86710.86960.92350.77980.5634

该表格可作为复现时的对照基线:读者在自己机器上跑同样的命令,若验证集指标与上表基本一致,说明环境与流程无误。

八、工程经验小结:把这套方案用到你自己的文本数据上

综合以上五个案例,可以沉淀出几条可迁移的实战经验:

  1. 统一入口 + 任务分发:把"数据加载、训练、预测、提交"封装成run_competition.py式的参数化脚本,新增比赛只需实现一个load_xxx()函数并注册到--task分发逻辑中。
  2. 优先尝试 stacking 模式:README 四个案例的主命令全部使用--mode stacking(5 折 Bagging + 单层 Stacking),将文本模型与表格模型互补融合;算力紧张时可降级为weighted甚至single
  3. 善于做标签变换:价格、销量等长尾回归目标先做log/log10变换再训练,提交前逆变换还原,这是三个回归案例共同的得分点。
  4. 文本特征工程:对半结构化文本(如 "123 out of 5 stars"、层级品类串)先做规则清洗与字段拆分,让预训练语言模型与表格模型都能吃到更干净的信号。
  5. 善用 preset 与 seed:通过--preset在训练速度与质量之间权衡,通过固定--seed 123保证实验可复现。
  6. 迁移与复用:GLUE 案例证明,一个训练好的文本模型可以直接复用到同分布的下游任务(AX 诊断集直接用 MNLI 检查点推理),这对实际业务中的冷启动非常有价值。

所有命令与脚本均可在仓库的 examples/automm/text_prediction 目录下找到,配合 multimodal/src/autogluon/multimodal/utils/presets.py 与 tabular/src/autogluon/tabular/configs/hyperparameter_configs.py 阅读源码,即可把这套比赛级文本预测方案移植到自己的数据集上。

【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/16 7:00:19

2026年9月 Java 面试题整理:200道 Java 高频面试题复习清单

2026年9月 Java 面试题整理:200道 Java 高频面试题复习清单 这份清单整理了 200道Java面试题,覆盖Java基础、集合、并发、JVM、Spring、MySQL、Redis、消息队列、分布式、系统设计和算法,适合面试前查漏补缺。同时纳入项目深挖、持久层、网络…

作者头像 李华
网站建设 2026/9/16 6:59:15

多模态大模型驱动数字孪生:从可视化到智能交互的实战

这几年做数字孪生项目,我最大的感受是:行业里不缺三维可视化能力,缺的是让这个“数字双胞胎”真正会思考、会交流的能力。传统的数字孪生系统,说到底就是把设备状态、传感器数据搬到屏幕上,人去看、去分析、去决策。但…

作者头像 李华
网站建设 2026/9/16 6:58:26

Win10系统安装全攻略:官方ISO与PE安装及UEFI/GPT分区匹配指南

微软原版ISO、PE安装、UEFI与Legacy分区这些关键词,几乎每个装系统的朋友都绕不开。今天这篇就一次性把这些说透,包括官方ISO直装和微PE两种方法,以及UEFIGPT和LegacyMBR这两种分区模式的区别和选择逻辑。如果你正准备自己重装Win10&#xff…

作者头像 李华
网站建设 2026/9/16 6:58:25

校园社团管理系统:SpringBoot+Vue3全栈开发实践

1. 项目概述:校园社团管理系统的技术架构与核心价值校园社团作为学生课外活动的重要载体,其管理效率直接影响着学生参与度和组织活力。传统基于Excel或纸质档案的管理方式存在信息孤岛、流程繁琐、数据易丢失等问题。这套基于Java SpringBootVue3MyBatis…

作者头像 李华
网站建设 2026/9/16 6:58:24

RDMA与GPUDirect:从内存旁路到GPU直接通信的技术解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/16 6:57:14

MatlabPNM:面向岩心CT数据的孔隙网络建模与渗流仿真框架

简介:本资源是一个面向科研人员与工程技术人员的Matlab孔隙网络建模工具包,专为多孔介质中流体流动、扩散传质及化学反应等过程的数值模拟而设计,适用于石油工程、地质学、环境科学和材料科学等领域。包内共33个文件,含14个核心功…

作者头像 李华