资讯动态

TabFM 三段式注意力架构完整拆解:列注意力、行注意力与ICL块如何协同工作

发布时间:2026/10/3 16:58:28 来源:尧图企业网站定制
TabFM 三段式注意力架构完整拆解列注意力、行注意力与ICL块如何协同工作【免费下载链接】tabfmTabFM (Tabular Foundation Model) is a pretrained tabular foundation model developed by Google Research for tabular data regression and classification.项目地址: https://gitcode.com/gh_mirrors/ta/tabfmTabFMTabular Foundation Model是 Google Research 开发的表格数据预训练基础模型无需在你的数据上训练即可通过上下文学习In-Context Learning对混合类型的表格数据做零样本分类与回归。本文带你完整拆解 TabFM 内部的三段式架构列注意力、行注意力与 ICL 块的协同原理以及 v1.0.0 的关键配置。一、为什么需要三段式注意力传统表格模型如 XGBoost把每行当作独立样本无法读懂整张表的统计规律而普通 Transformer 若直接对表格做注意力复杂度会随行列数平方级爆炸。TabFM 的解法是把列的分布规律、行内的特征交互、整表的监督信号拆成三级注意力逐级抽象数据流向总览v1.0.0 实际执行顺序阶段模块输入 → 输出作用① 单元格嵌入CellEmbedder每个单元格标量 → 256维向量Fourier 特征 特征分组编码② 列注意力ColEmbedding每列 T 个向量学习每列的分布规律③ 行注意力第一轮RowInteraction每行 H8 个向量行内特征交互保留全序列④ 列注意力第二轮ColEmbedding_2交叉信息再按列聚合弥补两轮行交互⑤ 行注意力第二轮RowInteraction_2每行 → 8 个 CLS 向量拼接压缩成行表示⑥ ICL 块ICLearning行表示 标签 y → 预测数据集级上下文学习核心源码全部位于 model.py文件头部注释定义了 B批、T行数、H列数、E嵌入维度等张量轴读源码前建议先扫一眼这份图例。二、第一段列注意力Column Attention——读懂每一列的分布2.1 单元格嵌入傅里叶特征 特征分组CellEmbedder 负责把表格里每个格子变成高维向量有三个巧妙设计傅里叶特征原始数值先乘以可学习的频率银行再取 sin/cos 拼接成 64 维32 频率 × 2送入线性层。这让模型获得丰富的频谱视角比裸线性投影更擅长刻画数值分布特征分组feature group以大小为 3 的滑窗对相邻列做重叠分组让每个单元格的编码天然包含邻列信息y 值融合v1.0.0 采用ADD_Y_TO_X_POST_EMBEDDING方案见 YEmbeddingScheme把标签 y 嵌入后只加到训练行的单元格上——这相当于把特征→标签的对应关系提前写进了嵌入里。2.2 Set TransformerO(n) 复杂度的诱导注意力列注意力的主体是 ColEmbedding内部使用 SetTransformer。它的核心技巧是 InducedSelfAttentionBlock第一阶段一组可学习的诱导点v1.0.0 为 256 个作为 Query去关注整列数据得到列的摘要向量第二阶段整列数据再反向关注这些摘要向量实现信息混合。这样注意力开销从 O(T²) 降为 O(T)100 行的上下文列和 10000 行列的计算量级相同这是 TabFM 能处理大表的关键。 列注意力是逐列独立处理的把 (B, T, H, E) 转置成 (B·H, T, E)每一列当作一个独立序列过 Set Transformer。三、第二段行注意力Row Attention——让一行内的特征互相说话RowInteraction 把视角从列切换到行捕捉行内特征交互比如面积 × 地段共同决定房价CLS 标记聚合每行序列头部拼接 8 个可学习 CLS tokenv1.0.0 配置row_num_cls8让它们像摘要官一样汇总该行所有特征RoPE 旋转位置编码行注意力编码器开启 RoPErope_base100000使模型感知特征在行内的列位置——第 3 列和第 17 列的含义可以不同两轮行注意力v1.0.0 实际跑了两轮列→行→列→行的交替见 TabFM.call第一轮保留全序列输出output_full_sequenceTrue供第二轮列注意力再聚合第二轮才压缩为 8 个 CLS 拼接的 2048 维行表示256×8送入 ICL 块。四、第三段ICL 块In-Context Learning——把训练集当上下文读ICLearning 是整张表的总解码器也是零样本能力的来源标签编码分类任务的 y 经OneHotAndLinear投影成 d_model 维向量回归任务的 y 经 MLP 编码加法注入标签向量加到行表示上仅训练行测试行的 y 被掩码清零24 层 Transformerv1.0.0 的 ICL 编码器有24 个块、8 个头列/行注意力各只有 3 块是算力最重的部分——它需要在带标签的训练行和无标签的测试行之间做全局注意力注意力掩码测试行只能关注训练行防止偷看未来的标签解码器输出过一个 MLP分类任务产出类别 logits支持最多 10 类回归任务产出连续值。整个前向流程可在 TabFM 主类 的类文档字符串中找到权威描述PyTorch 版同构实现在 tabfm/src/pytorch/model.py。五、Prefill / Decode 缓存像 LLM 一样推理TabFM 借鉴了大语言模型的推理模式见 prefill 与 decodePrefill一次性读取全部训练数据缓存三个模块的中间表示col1/col2诱导表示、icl的 KV CacheDecode测试数据只走轻量路径直接复用缓存预测耗时只与测试样本数相关与训练集大小无关。scikit-learn 封装层 TabFMClassifier 进一步提供了n_estimators集成机制对特征/行子采样、列置换生成 32 个成员投票进一步提升稳定性。六、v1.0.0 关键配置速查表所有硬编码参数集中在 Config参数值说明embed_dim256列/行注意力嵌入维度col_num_blocks / col_nhead3 / 4列注意力块数与头数col_num_inds256诱导点数row_num_blocks / row_nhead3 / 8行注意力块数与头数row_num_cls8每行 CLS token 数icl_num_blocks / icl_nhead24 / 8ICL 编码器规模max_classes10分类任务最多 10 类activationswiglu前馈激活函数use_fourier_featuresTrue启用傅里叶特征七、一分钟上手与常见限制本地快速安装JAX CPU 后端git clone https://gitcode.com/gh_mirrors/ta/tabfm cd tabfm pip install -e .[jax]然后运行 examples/classification_example.py 或 regression_example.py 即可体验零样本预测。新手最容易踩的四个限制 上下文行数默认 100max_num_rows控制每个集成成员读取的训练行数大表会被自动采样见 README FAQ特征数上限 500max_num_features默认 500超出的列会被子采样分类最多 10 类由max_classes10决定超出的数据集会直接报错权重许可源码是 Apache-2.0但预训练权重受tabfm-non-commercial-v1.0限制仅限非商用商用前务必确认。八、总结三段式架构的设计哲学一句话记忆 列注意力管分布行注意力管交互ICL 块管监督。列注意力用诱导点实现 O(n) 复杂度解决表太长问题行注意力用 CLS RoPE 把行内交互压缩成固定长度表示解决列太多问题ICL 块用 24 层深 Transformer 把带标签上下文整体读入用掩码隔离测试行解决不训练也能学的问题。三者接力完成单元格 → 列 → 行 → 预测的逐级抽象这正是 TabFM 能够零样本泛化到陌生表格数据的核心原因。想深入细节建议按 ColEmbedding → RowInteraction → ICLearning 的顺序精读源码配合 CHANGELOG 了解 v1.0.1 的推理缓存优化。【免费下载链接】tabfmTabFM (Tabular Foundation Model) is a pretrained tabular foundation model developed by Google Research for tabular data regression and classification.项目地址: https://gitcode.com/gh_mirrors/ta/tabfm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑