
基于 AutoGluon 与 TCGA 临床数据预测癌症患者生存状态TCGA-HNSC 实战与 FT-Transformer 深度解析【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon导读本文基于 AutoGluon 官方示例仓库中的 TCGA 癌症生存预测示例完整演示如何仅凭患者的临床信息年龄、性别、癌症分期、接受的治疗方式等构建二分类模型预测头颈鳞状细胞癌Head and Neck Squamous Cell Carcinoma, HNSC患者的生存状态。通过本文你将掌握TCGA 临床数据的获取与预处理要点、使用TabularPredictor一键训练多种模型的完整流程以及 AutoGluon 中 FT-TransformerFT_Transformer模型的底层实现原理、默认超参数与调优方法并能在自己的表格数据集上复现这套生存分析流水线。1. 任务背景与数据集本示例解决的问题是根据患者的临床记录预测患者是否存活二分类任务标签列为vital_status。数据来自美国癌症基因组图谱计划The Cancer Genome Atlas, TCGA的 TCGA-HNSC 项目该项目的临床信息以 TSV制表符分隔文件形式提供样本量约 1000 行、29 列属于典型的小规模高维表格数据。值得强调的是AutoGluon 在这里使用的是纯表格tabular建模管线而非多模态管线——脚本从autogluon.tabular导入TabularPredictor与TabularDataset其核心关注点是表格数据上的深度学习模型FT-Transformer。源码位于 example_cancer_survival.py。1.1 数据集下载与加载脚本内置了数据集的元信息通过 AutoGluon 多模态工具包中的download函数下载并校验 SHA1 哈希确保数据完整性INFO { name: cancer_survival.tsv, url: s3://automl-mm-bench/life-science/clinical.tsv, sha1sum: 6d19609c2a8492f767efd9f2c0b7687bcd3845a3 }加载逻辑位于data_loader函数example_cancer_survival.py若本地不存在cancer_survival.tsv则自动下载随后用pd.read_csv(full_path, sep\t)读取。1.2 临床数据预处理要点preprocess函数example_cancer_survival.py展示了医疗表格数据清洗的三个关键步骤这些做法对任何临床数据集都具备通用价值缺失值标记df df[df ! --]将数据集中的--占位符统一替换为NaN交由 AutoGluon 内置的缺失值处理逻辑处理删除无效列自动丢弃所有名称中含id的列如样本 ID以及全数据集取值唯一的常量列nunique() 1——这类列对模型没有任何判别信息剔除标签捷径shortcut列显式删除days_to_death距死亡天数、year_of_death死亡年份两列。这是生存类任务最容易踩的坑若保留距死亡天数这类列模型会学到死亡日期非空即已去世的捷径导致预测分数虚高、完全失真。训练目标既然是预测患者当前是否存活所有与死亡时间直接绑定的字段都必须移出特征集。完成清洗后通过sklearn.model_selection.train_test_split按test_size0.3、shuffleTrue划分训练集与测试集。2. 运行实验从单模型到全模型 Benchmark示例脚本通过命令行参数控制实验模式核心用法如下摘自 README.md# 在多个 AutoGluon 模型上做 Benchmark python3 example_cancer_survival.py --task TCGA_HNSC --mode all_models # 只跑 FT-Transformer python3 example_cancer_survival.py --task TCGA_HNSC --mode FT_Transformer2.1 命令行参数全解脚本的get_parser函数example_cancer_survival.py定义了以下参数参数类型默认值可选值 / 说明--pathstr./datasetTCGA 数据集保存目录--test_sizefloat0.3测试集占比划分时按该比例切分--shuffleboolTrue划分前是否打乱样本--seedint123随机种子同时作用于 PyTorch、NumPy 与 Pythonrandom--taskstradultTCGA_HNSC生存预测或adult成人收入预测--modestrall_modelsFT_Transformer仅深度模型或all_models完整模型库--num_gpusint-1GPU 数量-1表示自动探测全部可用 GPU--num_workersint2数据加载进程数2.2 训练配置解析train函数example_cancer_survival.py的核心逻辑如下metric accuracy hyperparameters {} if args.mode FT_Transformer else get_hyperparameter_config(default) hyperparameters[FT_TRANSFORMER] {env.num_gpus: args.num_gpus, env.num_workers: args.num_workers} predictor TabularPredictor(labellabel, eval_metricmetric).fit( train_datadf_train, hyperparametershyperparameters, time_limit900, )其中值得注意的细节get_hyperparameter_config(default)该函数定义于 hyperparameter_configs.py返回 AutoGluon 的默认模型组合包含NN_TORCH、三种不同规模/配置的GBMLightGBM含extra_trees变体与learning_rate0.03的大模型变体、CATCatBoost、XGB、FASTAIFastAI 神经网络、两种随机森林RFGini/Entropy 准则与两种极端随机树XTmodeFT_Transformer时hyperparameters{}空字典意味着跳过所有默认模型随后只注入FT_TRANSFORMER键从而单独训练 FT-Transformertime_limit900所有模型共享 900 秒15 分钟训练时间预算指标为accuracy二分类准确率同时用于模型选择与早停判断训练完成后调用predictor.leaderboard(df_test)生成测试集排行榜并保存到./leaderboard.csv。FT-Transformer 在 TabularPredictor 中的注册名为FT_TRANSFORMER其封装类FTTransformerModel位于 tabular/src/autogluon/tabular/models/automm/ft_transformer.py属于MultiModalPredictorModelautomm_model.py的派生实现内部实际驱动的是一个autogluon.multimodal.MultiModalPredictor。从源码看它与普通 AutoMM 封装模型的关键差异在于minimum_num_gpus 0、gpu_required False即允许纯 CPU 训练虽然脚本中仍强烈建议使用 GPU见下方警告逻辑valid_raw_types[R_INT, R_FLOAT, R_CATEGORY]即只处理数值与类别特征不处理文本集成阶段默认fold_fitting_strategy_gpusequential_local因为并行 bagging 在 GPU 上会引发崩溃。在 CPU 上训练时_fit方法会输出明确警告ft_transformer.py「未指定 GPU训练可能耗时很长建议使用 GPU 加速」。3. FT-Transformer 源码级原理剖析FT-TransformerFeature Tokenizer Transformer最初由 Gorishniy 等人在Revisiting Deep Learning Models for Tabular DataNeurIPS 2021中提出。AutoGluon 在 multimodal/src/autogluon/multimodal/models/ft_transformer.py 中提供了完整的 PyTorch 实现可将其拆解为三层结构。3.1 特征分词器Feature TokenizerFT-Transformer 的核心思想是把表格中的每个特征当作一个词用嵌入向量表示后送入 Transformer。AutoGluon 实现了两类分词器CategoricalFeatureTokenizerft_transformer.py为每个类别特征维护独立的nn.Embedding表通过category_offsets将各列类别索引偏移后查表得到 token每个特征还可附加一个不共享的可训练 bias 向量token_biasTrue时使 token 蕴含该特征取任何值都存在的位置信息NumEmbeddingsft_transformer.py数值特征先经分箱piecewise linear encoding再映射为嵌入从而摆脱传统 MLP 对数值特征单一标量输入的容量瓶颈。3.2 Transformer 骨干与分类头FT_Transformer类ft_transformer.py负责将各特征的 token 拼接mergeconcat在序列首部附加一个[CLS]类型的聚合 tokenpooling_modecls然后送入多层 Transformer 编码器。编码器采用预归一化prenormalization、GEGLU 激活的前馈网络ffn_activationgeglu最终分类头使用 MLP 输出预测概率。3.3 默认架构配置FT-Transformer 的默认架构在 configs/model/default.yaml 中定义ft_transformer: data_types: [categorical, numerical] embedding_arch: [linear] token_dim: 192 # 每个特征 token 的维度 hidden_size: 192 # Transformer 骨干的嵌入维度 num_blocks: 3 # Transformer 编码器层数 attention_num_heads: 8 # 多头注意力头数 attention_dropout: 0.2 residual_dropout: 0.0 ffn_dropout: 0.1 ffn_hidden_size: 192 ffn_activation: geglu head_activation: relu normalization: layer_norm merge: concat pooling_mode: cls # 用 [CLS] token 聚合序列 checkpoint_name: null从配置文件可以看出默认模型为 3 层 Transformer、隐藏维度 192、8 注意力头属于轻量级配置可同时处理类别与数值特征。3.4 训练期默认超参数封装层FTTransformerModel._set_default_paramstabular/src/autogluon/tabular/models/automm/ft_transformer.py进一步固化了训练期行为理解这些默认值有助于解释为何 FT-Transformer 训练较慢、效果却更稳参数默认值含义与影响model.ft_transformer.embedding_arch[linear]数值特征采用线性嵌入架构env.batch_size/env.per_gpu_batch_size128/128训练 batch 大小optim.max_epochs2000上限极大配合早停实现训练至收敛optim.weight_decay1.0e-5权重衰减抑制过拟合optim.lr_schedulepolynomial_decay多项式衰减学习率调度optim.patience20验证指标连续 20 个 epoch 无提升则早停optim.top_k3保留 top-k 检查点再择优_max_features300特征数上限超过则直接跳过该模型见 automm_model.py其中_max_features300是一个值得注意的约束若输入特征数超过 300FT-Transformer 会抛出AssertionError并跳过训练开发者注释中也明确标注这是hack未来版本可能改为正式的ag_args_fit参数。对于特征列极多100 列的数据集官方在 predictor.py 中同样提示 FT-Transformer 不擅长扩展到超过 100 个特征建议改用 TabM 类模型。4. 实验结果TCGA-HNSC 与 adult 双数据集对照所有模型以 900 秒为时间上限训练结果如下完整排行榜保存于./results/leaderboard.csv。4.1 TCGA-HNSC约 1000 行 × 29 列测试集 30%模型 (TCGA-HNSC)测试准确率验证准确率训练时间 (s)测试时间 (s)NeuralNetTorch0.9432180.9256762.7002170.027071RandomForestGini0.9400630.8918920.6038930.108412LightGBMLarge0.9085170.9391891.3510580.014151CatBoost0.9053630.9391894.8044890.025413XGBoost0.9053630.9256760.4168470.027664WeightedEnsemble_L20.8738170.9459461.6368080.028049FTTransformer0.8643530.89189251.3058470.3846694.2 adult49K 样本更大规模对照为验证模型在更大数据规模下的表现可将--task切换为adultUCI 成人收入数据集约 49000 条实例标签列class训练数据与测试数据由脚本分别从 AutoGluon 公共 S3 自动下载example_cancer_survival.pypython3 example_cancer_survival.py --task adult --mode all_models模型 (adult)测试准确率验证准确率训练时间 (s)测试时间 (s)XGBoost0.8771620.88720.6979710.038446WeightedEnsemble_L20.8765480.890842.2010000.316964CatBoost0.8748080.88284.2632310.016138FTTransformer0.8622170.8696820.4858782.732730RandomForestEntr0.8579180.86200.9798200.249948NeuralNetFastAI0.8573040.862032.1061480.137593NeuralNetTorch0.8563820.858840.2640390.1770794.3 结论解读从两组结果可以得出三个客观结论注意这些数据仅反映示例脚本在特定时间预算下的单次运行结果不代表模型的全量性能上限梯度提升树GBDT家族依然领跑XGBoost、CatBoost、LightGBM 在两个数据集上均处于第一梯队且训练速度极快这与 AutoGluon 在表格任务上树模型通常是强基线的既有经验一致FT-Transformer 是深度模型中的最佳选择在 adult 数据集上FT-Transformer 的测试准确率0.862217全面超越 NeuralNetFastAI 与 NeuralNetTorch 等传统深度模型——README 中明确指出虽然决策树类模型仍是最优方法但 FT-Transformer 击败了其他深度学习方法训练成本差距显著FT-Transformer 在 adult 上耗用了约 820 秒接近全部 900 秒预算而树模型仅需数秒。这与其逐 epoch 迭代 早停的深度训练范式直接相关也解释了为何在time_limit900的约束下它的成绩仍略逊于树模型——它是用时间换深度的候选模型适合放在 ensemble 中补充多样性。5. 如何把示例迁移到自己的表格数据集该脚本的train函数已内置了两种任务的接入模式example_cancer_survival.py迁移到自有数据只需三处改动数据接入将df_train、df_test替换为自己的 DataFrame可通过TabularDataset从文件/URL 加载或直接传入pandas.DataFrame标签列将label变量设为目标列名若为多分类或回归问题无需改动脚本TabularPredictor会自动推断问题类型评估指标metric accuracy可替换为 AutoGluon 支持的其他指标如roc_auc、f1、log_loss完整指标表可参考 core 模块的 metrics 实现与对应测试 test_classification_metrics.py。需要特别注意的移植陷阱仍与第 1.2 节一致务必剔除与标签存在时序泄漏的捷径列如事件发生日期、结果相关的时间戳否则验证分数会严重虚高。6. 小结本文以 examples/automm/TCGA_cancer_survival 为骨架完整还原了 AutoGluon 在 TCGA-HNSC 临床数据上的生存状态预测流水线从带 SHA1 校验的数据下载、三步数据清洗缺失值、id/常量列、捷径列到--task/--mode/--num_gpus等参数的实验编排再到 FT-Transformer 的源码级剖析与双数据集基准对照。实践层面你可以直接复用该脚本验证 AutoGluon 在医疗表格数据上的开箱即用能力理论层面通过 ft_transformer.py 与 ft_transformer.pytabular 封装 的对照阅读也能深入理解 AutoGluon 如何把多模态框架的 Transformer 能力降维复用到纯表格任务并借助集成学习获得稳健的预测表现。【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考