表格神经网络架构发展史:从 MLP 到表格基础模型的二十年

表格神经网络架构发展史:从 MLP 到表格基础模型的二十年

导语 :表格数据(tabular data)------行样本、列特征的结构化表格------是现实世界最常见的数据形态,据估计占据企业数据 70% 以上。然而与图像、文本不同,表格长期是深度学习的"滑铁卢":决策树集成(GBDT / XGBoost)在大多数基准上稳压神经网络。过去近二十年,学术界围绕"如何让神经网络真正吃好表格数据"演化出了一条清晰的架构演进脉络。本文按五个阶段系统梳理,每个代表架构均附论文原图示例代码、设计说明、实例与得失分析。

目录

  • 背景:为什么表格数据对神经网络如此困难
  • [第一阶段:MLP 与推荐双塔时代(~2016)](#第一阶段:MLP 与推荐双塔时代(~2016))
    • [1.1 朴素 MLP](#1.1 朴素 MLP)
    • [1.2 Wide & Deep(2016)](#1.2 Wide & Deep(2016))
    • [1.3 DeepFM(2017)](#1.3 DeepFM(2017))
  • 第二阶段:注意力机制引入(2019-2020)
    • [2.1 TabNet(2019)](#2.1 TabNet(2019))
    • [2.2 TabTransformer(2020)](#2.2 TabTransformer(2020))
  • [第三阶段:Transformer 全面进军(2021)](#第三阶段:Transformer 全面进军(2021))
    • [3.1 FT-Transformer(2021)](#3.1 FT-Transformer(2021))
    • [3.2 SAINT(2021)](#3.2 SAINT(2021))
    • [3.3 阶段评价](#3.3 阶段评价)
  • [第四阶段:简约之美------ResNet 基线的反思(2021-2023)](#第四阶段:简约之美——ResNet 基线的反思(2021-2023))
    • [4.1 ResNet-MLP 强基线](#4.1 ResNet-MLP 强基线)
    • [4.2 TabR(2023)](#4.2 TabR(2023))
  • 第五阶段:表格基础模型与上下文学习(2022-2026)
    • [5.1 TabPFN(2022)](#5.1 TabPFN(2022))
    • [5.2 TabICL 与检索增强(2024-2026)](#5.2 TabICL 与检索增强(2024-2026))
    • [5.3 TabFM 与商业基础模型(2024-2025)](#5.3 TabFM 与商业基础模型(2024-2025))
    • [5.4 LLM 与表格理解(2024-2026)](#5.4 LLM 与表格理解(2024-2026))
    • [5.5 三大表格基础模型对比:TabPFN vs TabICL vs TabFM](#5.5 三大表格基础模型对比:TabPFN vs TabICL vs TabFM)
  • 全景总结:五条主线,一条螺旋
  • 给从业者的选型建议
  • 深度问答:两个核心疑问的澄清
    • [Q1:为什么 MLP 无法自动发现高阶特征交互?加深网络行不行?](#Q1:为什么 MLP 无法自动发现高阶特征交互?加深网络行不行?)
    • [Q2:为什么 FT-Transformer 不能像 LLM 一样训一次、处处迁移?](#Q2:为什么 FT-Transformer 不能像 LLM 一样训一次、处处迁移?)
  • 参考来源

背景

表格数据之所以对神经网络如此"难啃",根源在于三个本质难点,它们贯穿了整条演进线:

  1. 特征异构性强:一列是年龄(连续),下一列是邮政编码(高基数类别),再下一列是布尔。CNN 的卷积核、RNN 的时序假设在这里统统失效,无法像图像那样在局部共享参数。
  2. 无空间/时序结构:列的顺序对预测通常无意义,换列不影响语义。这打破了"位置即信息"的归纳偏置------Transformer 把每个特征 token 化才绕过了这一点。
  3. 数据量小、噪声大:很多表格任务样本量只有几千到几万,深度模型极易过拟合;而树模型自带特征选择与稀疏划分,天然契合这类场景。

正因如此,早期 NN 架构不得不从"如何处理异构特征""如何建模特征交互"这两个最基础的问题开始,一步步演化。

代码约定 :下文所有示例代码均基于 PyTorch,仅作架构示意,省略训练循环与超参细节。重点展示特征处理方式网络结构


第一阶段

1.1 朴素 MLP

最早的表格神经网络就是多层感知机(MLP):把所有特征拼接成一个向量,过几层全连接 + ReLU,输出预测。看似简陋,却是后续所有架构的对照基线。

架构示意

复制代码
特征向量 x (连续列标准化 + 类别列 one-hot 拼接)
    → 全连接 + ReLU (dim=256)
    → 全连接 + ReLU (dim=128)
    → 全连接 + ReLU (dim=64)
    → 输出层 (回归/分类)

示例代码

python 复制代码
import torch
import torch.nn as nn

class TabularMLP(nn.Module):
    def __init__(self, num_continuous, cat_cardinalities, hidden_dims, d_out=1):
        super().__init__()
        # 类别特征:每个类别列一个 embedding 表
        self.cat_embeddings = nn.ModuleList([
            nn.Embedding(card, min(16, (card + 1) // 2))  # 简单的维度选择
            for card in cat_cardinalities
        ])
        cat_dim = sum(min(16, (c + 1) // 2) for c in cat_cardinalities)

        # 连续特征直接使用(标准化后)
        self.num_bn = nn.BatchNorm1d(num_continuous)

        # MLP 主体
        in_dim = num_continuous + cat_dim
        layers = []
        for h in hidden_dims:
            layers += [nn.Linear(in_dim, h), nn.ReLU(), nn.BatchNorm1d(h)]
            in_dim = h
        self.mlp = nn.Sequential(*layers)
        self.head = nn.Linear(in_dim, d_out)

    def forward(self, x_num, x_cat):
        # 特征处理:连续标准化 + 类别嵌入查表
        x_num = self.num_bn(x_num)                           # [B, num_cont]
        cat_embs = [emb(x_cat[:, i]) for i, emb in
                     enumerate(self.cat_embeddings)]        # 每个 [B, emb_dim]
        x = torch.cat([x_num] + cat_embs, dim=1)             # [B, in_dim]
        return self.head(self.mlp(x))                        # [B, d_out]

# 使用示例:3 个连续特征 + 2 个类别特征(城市 50 类、职业 20 类)
model = TabularMLP(
    num_continuous=3,          # 年龄、收入、负债比
    cat_cardinalities=[50, 20],  # 城市、职业
    hidden_dims=[256, 128, 64],
    d_out=1                   # 二分类(违约概率)
)
# 输入:x_num=[B,3], x_cat=[B,2](整数索引)
# 输出:[B,1] 预测值

举例:房价预测------输入 8 维特征(面积、卧室数、房龄、地段编码等),3 层 256-128-64 全连接,输出房价。它对所有特征一视同仁地线性组合,无法自动发现"地段 × 面积"这类高阶交互,需要依赖人工构造交叉列。

解决了什么问题:第一次给出表格任务的端到端可微模型,可统一处理连续与类别特征,是深度学习进入表格的最小可行方案。

效果:在数据量充足、特征已做好工程化时可用;但小数据上易过拟合,且强依赖人工特征工程,绝大多数场景被 GBDT 压制。Gorishniy 等人 2021 年的实验也表明,调参得当的 MLP + 适当预处理其实不弱------这个"返璞归真"结论后来反复出现,成为悬在所有新架构头上的达摩克利斯之剑。

优化点

  • 数值与类别特征混合拼接,表征不统一,类别信息被稀释;
  • 缺乏自动特征交互,全靠人工;
  • 无正则/归一化设计,小数据不稳。

1.2 Wide & Deep

2016 年 Google 提出 Wide & Deep ,真正让表格 NN 走出实验室、进入工业推荐场景。核心思想是显式与隐式特征交互的分工

论文原图 (Wide & Deep 架构,来源:wngaw.github.io):

示例代码

python 复制代码
import torch
import torch.nn as nn

class WideAndDeep(nn.Module):
    def __init__(self, num_continuous, cat_cardinalities,
                 cross_features_dim, embed_dim=8, hidden_dims=[128, 64], d_out=1):
        super().__init__()
        # --- Deep 部分:类别特征 → embedding ---
        self.cat_embeddings = nn.ModuleList([
            nn.Embedding(card, embed_dim) for card in cat_cardinalities
        ])
        deep_in = num_continuous + len(cat_cardinalities) * embed_dim

        layers = []
        for h in hidden_dims:
            layers += [nn.Linear(deep_in, h), nn.ReLU()]
            deep_in = h
        self.deep = nn.Sequential(*layers)

        # --- Wide 部分:直接使用手工交叉特征(线性) ---
        self.wide = nn.Linear(cross_features_dim, 1, bias=False)

        # 融合输出
        self.head = nn.Linear(hidden_dims[-1] + 1, d_out)

    def forward(self, x_num, x_cat, x_cross):
        # Deep: embedding + MLP
        cat_embs = [emb(x_cat[:, i]) for i, emb in enumerate(self.cat_embeddings)]
        deep_in = torch.cat([x_num] + cat_embs, dim=1)
        deep_out = self.deep(deep_in)                       # [B, hidden]

        # Wide: 线性模型直接吃交叉特征
        wide_out = self.wide(x_cross)                       # [B, 1]

        # 融合
        return self.head(torch.cat([deep_out, wide_out], dim=1))

# 使用示例:
model = WideAndDeep(
    num_continuous=5,            # 年龄、收入等连续特征
    cat_cardinalities=[100, 50], # user_id, item_id
    cross_features_dim=20,       # 手工构造的交叉特征维度
    d_out=1                     # CTR 二分类
)
# 输入:x_num=[B,5], x_cat=[B,2], x_cross=[B,20]

举例 :应用商店 App 推荐中,Wide 部分输入手工交叉特征如 gender=男 AND age_bucket=20-30 AND app_category=游戏,负责记住"这个人群爱下游戏"这类强信号;Deep 部分从 embedding 学到未见过的组合,负责向新 App 泛化。

解决了什么问题:单一 MLP 难以同时兼顾"记忆"(强规则、长尾共现)与"泛化"(新组合),Wide & Deep 用双塔结构把两者显式分工,既不丢强记忆信号,又能泛化。

效果:在 Google Play 推荐上线后,App 获得率显著提升,成为推荐系统深度化落地的里程碑。

优化点

  • Wide 部分仍依赖人工构造交叉特征,工程成本高;
  • 两路嵌入不共享,参数效率低;
  • 仅二阶以内的线性交叉,高阶交互靠 Deep 隐式学,可控性差。

1.3 DeepFM

2017 年华为提出 DeepFM ,把 Wide & Deep 的手工交叉换成了 FM(Factorization Machine),实现去人工化的特征交互

论文原图 (DeepFM 架构,来源:arXiv:1703.04247):

示例代码

python 复制代码
import torch
import torch.nn as nn

class DeepFM(nn.Module):
    def __init__(self, num_continuous, cat_cardinalities,
                 embed_dim=8, hidden_dims=[128, 64], d_out=1):
        super().__init__()
        # 共享 embedding 层:FM 和 Deep 共用同一套嵌入
        self.embeddings = nn.ModuleList([
            nn.Embedding(card, embed_dim) for card in cat_cardinalities
        ])
        # 连续特征也映射到 embed_dim(线性投影)
        self.num_proj = nn.Linear(num_continuous, embed_dim)

        # --- FM 部分:一阶 + 二阶交叉 ---
        self.fm_first = nn.Linear(num_continuous + len(cat_cardinalities), 1)

        # --- Deep 部分:MLP ---
        deep_in = (len(cat_cardinalities) + 1) * embed_dim  # 类别嵌入 + 连续投影
        layers = []
        for h in hidden_dims:
            layers += [nn.Linear(deep_in, h), nn.ReLU()]
            deep_in = h
        self.deep = nn.Sequential(*layers)
        self.head = nn.Linear(hidden_dims[-1], d_out)

    def forward(self, x_num, x_cat):
        B = x_num.shape[0]
        # 共享嵌入
        cat_embs = [emb(x_cat[:, i]) for i, emb in enumerate(self.embeddings)]
        num_emb = self.num_proj(x_num)                       # [B, embed_dim]
        all_embs = torch.stack(cat_embs + [num_emb], dim=1)  # [B, n_fields, embed_dim]

        # FM 一阶项
        one_hot_like = torch.cat([x_num, x_cat.float()], dim=1)
        fm_1st = self.fm_first(one_hot_like)                  # [B, 1]

        # FM 二阶项:嵌入两两内积之和
        sum_sq = all_embs.sum(dim=1).pow(2).sum(dim=1)        # (∑ e_i)^2
        sq_sum = all_embs.pow(2).sum(dim=1).sum(dim=1)        # ∑(e_i^2)
        fm_2nd = 0.5 * (sum_sq - sq_sum).unsqueeze(1)        # [B, 1]

        # Deep 部分
        deep_in = all_embs.reshape(B, -1)
        deep_out = self.deep(deep_in)

        return self.head(deep_out) + fm_1st + fm_2nd         # 加法融合

# 使用示例:
model = DeepFM(
    num_continuous=5, cat_cardinalities=[1000, 500, 50],
    embed_dim=8, hidden_dims=[128, 64], d_out=1
)
# 输入:x_num=[B,5], x_cat=[B,3] → 自动学二阶交叉,无需手工构造

举例 :电商 CTR 预估中,FM 部分自动学习 user_id 嵌入 × item_id 嵌入 的二阶交互,无需人工写 user_id AND item_id 交叉列;Deep 部分 MLP 学高阶非线性关系。两者共享同一套 embedding,端到端训练。

解决了什么问题:消除了 Wide & Deep 的手工特征工程------FM 通过嵌入内积自动学二阶交叉,同时与 MLP 共享嵌入,参数更省、训练更简单。

效果:在多个 CTR 公开数据集上超越 Wide & Deep,成为推荐领域事实标准之一,后续衍生出 xDeepFM、AutoInt 等细化高阶交互的变体。

优化点

  • FM 仅建模二阶交叉,更高阶交互仍依赖 Deep 隐式学习;
  • 类别嵌入对低频字段学不充分(冷启动问题);
  • 连续特征的处理仍偏简单(投影到 embed_dim),异构性未完全解决。

第一阶段小结

这一代以推荐系统为主战场,核心命题是"如何把特征交互塞进网络"。Wide & Deep 和 DeepFM 让 NN 在 CTR 这类亿级样本场景站稳脚跟,但在通用表格基准(特征异构、样本数千~数万)上,GBDT 仍是统治者。下一阶段,注意力机制登场,试图解决"特征异构 + 自动选择"。


第二阶段

2.1 TabNet

2019 年 Google Cloud AI 提出 TabNet (Arik & Pfister),是第一个被广泛引用的"纯表格"深度架构。核心是逐步注意力(step-by-step attention),带有明显树模型启发------试图把"稀疏特征选择"融入神经网络。

论文原图 (TabNet 编码器架构,来源:arXiv:1908.07442):

图中展示了 TabNet 编码器的多步决策结构:每一步通过稀疏注意力选择部分特征,聚合后输出决策片段,多步求和得到最终预测。attention mask 直接对应特征重要度,提供内在可解释性。

示例代码

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

class SparseAttention(nn.Module):
    """稀疏门控:选部分特征进入下一步"""
    def __init__(self, input_dim, feat_dim):
        super().__init__()
        self.gate = nn.Linear(input_dim, feat_dim)

    def forward(self, h, prior):
        # h: 上一步的隐藏状态; prior: 累积 mask(避免重复选同一特征)
        logits = self.gate(h)                    # [B, feat_dim]
        # 乘以 (1 - prior) 实现"用过的特征降权"
        mask_logits = logits * (1.0 - prior)
        # 稀疏化:top-k 或 soft 稀疏
        mask = F.relu(torch.tanh(mask_logits))   # 软稀疏 mask
        return mask                              # [B, feat_dim]

class TabNetStep(nn.Module):
    """单步:特征选择 → 特征处理 → 决策输出"""
    def __init__(self, feat_dim, hidden_dim):
        super().__init__()
        self.attention = SparseAttention(hidden_dim, feat_dim)
        self.fc_feat = nn.Linear(feat_dim, hidden_dim)    # 处理选中特征
        self.fc_split = nn.Linear(hidden_dim, hidden_dim) # 分裂:一半给决策,一半传下一步

    def forward(self, x, h, prior):
        mask = self.attention(h, prior)           # [B, feat_dim] 稀疏 mask
        masked_x = x * mask                       # 选特征
        feat = F.gelu(self.fc_feat(masked_x))     # 处理
        split = self.fc_split(feat)
        decision_out, h_next = split.chunk(2, dim=1)  # 决策 + 传递
        return decision_out, h_next, mask

class TabNetEncoder(nn.Module):
    def __init__(self, feat_dim, hidden_dim=64, n_steps=3):
        super().__init__()
        self.n_steps = n_steps
        self.step0 = TabNetStep(feat_dim, hidden_dim)
        self.steps = nn.ModuleList([
            TabNetStep(feat_dim, hidden_dim) for _ in range(n_steps - 1)
        ])
        self.head = nn.Linear(hidden_dim // 2, 1)

    def forward(self, x):
        # x: [B, feat_dim] --- 所有特征拼接(连续标准化 + 类别嵌入后)
        B = x.shape[0]
        prior = torch.zeros_like(x)               # 累积 mask

        # 第一步
        h = x  # 简化:用输入初始化
        decision, h, mask = self.step0(x, h, prior)
        prior = prior + mask
        total_decision = decision

        for step in self.steps:
            decision, h, mask = step(x, h, prior)
            prior = prior + mask
            total_decision = total_decision + decision

        return self.head(total_decision), prior   # 预测 + 特征重要度

# 使用示例:
model = TabNetEncoder(feat_dim=20, hidden_dim=64, n_steps=3)
# 输入:x=[B, 20](已预处理的特征向量)
# 输出:预测值 + 特征重要度 mask(可解释性)

举例:信贷风控中,TabNet 第一步可能聚焦"收入 + 负债",第二步聚焦"历史逾期次数",最后预测违约概率,并能输出"哪些特征在决策中起了作用"。这种多步选择类似于决策树逐层判断。

解决了什么问题 :既想要端到端表格学习,又想要树模型式的特征选择 + 可解释性。TabNet 用稀疏 attention mask 实现逐特征、逐步骤的选择性输入,把"看哪些列"本身变成可学习的决策。

效果 :提供了内在可解释性(attention mask 直接对应特征重要度),在部分领域(如医疗、金融风控)有落地;但在多数学术基准上未稳定超过 GBDT,且多步结构训练慢、超参敏感。

优化点

  • 多步稀疏注意力训练开销大,收敛慢;
  • 稀疏 mask 的梯度估计不稳,特征选择不总是锐利;
  • 对连续特征的嵌入处理仍简单,异构性未根本解决;
  • 可解释性是"事后注意力",与真实因果贡献未必一致。

2.2 TabTransformer

2020 年 Hu et al. 提出 TabTransformer ,专攻一个被前人忽视的关键点------类别特征的语义嵌入

论文原图 (TabTransformer 架构,来源:arXiv:2012.06678):

图中展示了 TabTransformer 的核心设计:类别列经 embedding 后送入 Transformer encoder(多头自注意力),生成"上下文化嵌入";连续列则在 Transformer 之外简单拼接后过 MLP。

示例代码

python 复制代码
import torch
import torch.nn as nn

class TabTransformer(nn.Module):
    def __init__(self, num_continuous, cat_cardinalities,
                 embed_dim=32, n_heads=4, n_layers=2,
                 ffn_dim=64, mlp_dims=[128, 64], d_out=1):
        super().__init__()
        # 类别特征 → embedding
        self.cat_embeddings = nn.ModuleList([
            nn.Embedding(card, embed_dim) for card in cat_cardinalities
        ])
        # 连续特征 → 简单逐列投影到 embed_dim(原论文直接拼接,这里做投影便于理解)
        self.num_proj = nn.Linear(1, embed_dim)  # 每个连续列独立投影

        # Transformer encoder(只处理类别嵌入)
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embed_dim, nhead=n_heads,
            dim_feedforward=ffn_dim, batch_first=True
        )
        self.transformer = nn.TransformerEncoder(encoder_layer, n_layers)

        # MLP 头:处理 Transformer 输出 + 连续特征
        cat_out_dim = len(cat_cardinalities) * embed_dim
        mlp_in = cat_out_dim + num_continuous
        layers = []
        for h in mlp_dims:
            layers += [nn.Linear(mlp_in, h), nn.ReLU()]
            mlp_in = h
        self.mlp = nn.Sequential(*layers)
        self.head = nn.Linear(mlp_in, d_out)

    def forward(self, x_num, x_cat):
        # 类别 → embedding → Transformer
        cat_embs = [emb(x_cat[:, i]) for i, emb in
                     enumerate(self.cat_embeddings)]      # 每个 [B, embed_dim]
        cat_seq = torch.stack(cat_embs, dim=1)            # [B, n_cat, embed_dim]
        cat_out = self.transformer(cat_seq)               # [B, n_cat, embed_dim]
        cat_flat = cat_out.reshape(cat_out.shape[0], -1)   # [B, n_cat * embed_dim]

        # 连续特征:直接拼接(不经过 Transformer ------ 这是 TabTransformer 的短板!)
        x = torch.cat([cat_flat, x_num], dim=1)
        return self.head(self.mlp(x))

# 使用示例:
model = TabTransformer(
    num_continuous=5, cat_cardinalities=[50, 20, 10],
    embed_dim=32, n_heads=4, n_layers=2, d_out=1
)
# 输入:x_num=[B,5], x_cat=[B,3]
# 注意:连续特征只经过简单拼接,不参与自注意力建模!

举例:在 census 收入预测中,"职业=教师"和"职业=讲师"初始 embedding 可能相近,但经过 Transformer 自注意力后,它们会被同一行内的其他特征(如"行业=教育")上下文化调整,变得更贴合该样本语境。

解决了什么问题 :把 NLP 的"语境建模"迁移到表格------让类别嵌入不再是孤立的查表,而是在行内上下文中相互校准,缓解了深度模型的语义消歧能力不足。

效果:在类别列占主导、且列间有语义关联的任务上有效(如 census、森林覆盖类型);但把连续特征排除在 Transformer 之外简单拼接,成为它最大的短板------连续特征间的交互、连续与类别间的交互都被忽略。

优化点

  • 连续特征被"二等公民"对待,未进入自注意力建模;
  • 对高基数类别列的 embedding 学不充分(冷启动);
  • 无样本间信息流动,仍逐样本独立处理。

第二阶段小结

注意力机制登场,但两条路线各有遗憾:TabNet 追求稀疏选择与可解释性,却牺牲了效率;TabTransformer 让类别嵌入语义化,却冷落了连续特征。真正的突破要等下一阶段------把所有特征统一 token 化。


第三阶段

3.1 FT-Transformer

2021 年 Gorishniy 等人提出 FT-Transformer(Feature Tokenizer + Transformer),彻底解决了 TabTransformer 的连续特征短板,确立了"统一 token 化"范式。

论文原图 (FT-Transformer 架构,来源:arXiv:2106.11959):

图中展示了 FT-Transformer 的核心设计:Feature Tokenizer 将每一列(无论连续还是类别)统一映射为一个 token 向量,所有 token 送入标准 Transformer encoder,CLS token 做最终预测。

示例代码

python 复制代码
import torch
import torch.nn as nn

class FeatureTokenizer(nn.Module):
    """将每个特征(连续/类别)统一映射为 token 向量"""
    def __init__(self, num_continuous, cat_cardinalities, embed_dim=192):
        super().__init__()
        # 连续特征:标量 × 可学习权重 + bias → 向量
        self.num_tokenizer = nn.ModuleList([
            nn.Linear(1, embed_dim) for _ in range(num_continuous)
        ])
        # 类别特征:embedding 查表
        self.cat_embeddings = nn.ModuleList([
            nn.Embedding(card, embed_dim) for card in cat_cardinalities
        ])
        # [CLS] token
        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))

    def forward(self, x_num, x_cat):
        B = x_num.shape[0]
        # 连续特征 token 化:每列一个 token
        num_tokens = [
            proj(x_num[:, i:i+1]) for i, proj in enumerate(self.num_tokenizer)
        ]  # 每个返回 [B, embed_dim]
        # 类别特征 token 化:embedding 查表
        cat_tokens = [
            emb(x_cat[:, i]) for i, emb in enumerate(self.cat_embeddings)
        ]
        # 拼接所有 token
        tokens = torch.stack(num_tokens + cat_tokens, dim=1)  # [B, n_features, embed_dim]
        # 加 [CLS] token
        cls = self.cls_token.expand(B, -1, -1)               # [B, 1, embed_dim]
        tokens = torch.cat([cls, tokens], dim=1)             # [B, 1+n_features, embed_dim]
        return tokens

class FTTransformer(nn.Module):
    def __init__(self, num_continuous, cat_cardinalities,
                 embed_dim=192, n_heads=8, n_layers=3,
                 ffn_mult=4, d_out=1):
        super().__init__()
        self.tokenizer = FeatureTokenizer(num_continuous,
                                          cat_cardinalities, embed_dim)
        # 标准 Transformer encoder
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embed_dim, nhead=n_heads,
            dim_feedforward=embed_dim * ffn_mult,
            batch_first=True, activation='gelu'
        )
        self.transformer = nn.TransformerEncoder(encoder_layer, n_layers)
        # 用 [CLS] token 的输出做预测
        self.norm = nn.LayerNorm(embed_dim)
        self.head = nn.Linear(embed_dim, d_out)

    def forward(self, x_num, x_cat):
        tokens = self.tokenizer(x_num, x_cat)   # [B, 1+n, embed_dim]
        out = self.transformer(tokens)          # [B, 1+n, embed_dim]
        cls_out = self.norm(out[:, 0])          # 取 [CLS] token
        return self.head(cls_out)               # [B, d_out]

# 使用示例:
model = FTTransformer(
    num_continuous=5, cat_cardinalities=[50, 20, 10],
    embed_dim=192, n_heads=8, n_layers=3, d_out=1
)
# 输入:x_num=[B,5], x_cat=[B,3]
# 关键:连续和类别特征同等对待,都变成 token 送入 Transformer

举例:在银行客户流失预测中,"年龄=35"经 Feature Tokenizer 变成一个 192 维向量 token,"产品类型=信用卡"经 embedding 也是一个 192 维 token,两者与 CLS token 一起送入 Transformer。自注意力让"年龄"token 能直接与"产品类型"token 交互,不再有连续/类别的"二等公民"之分。

解决了什么问题 :统一了异构特征的表征------每一列无论连续还是类别,都被映射成同等维度的 token,送入同一 Transformer。消除了 TabTransformer 的连续/类别割裂,让所有特征在统一的注意力框架内交互。

效果:在 LAMDA-Tabular 综合基准中,FT-Transformer 在 Transformer 类方法里常胜出,成为后续最常被对照的强基线。在中等规模数据上能与 GBDT 持平或偶尔超越。

优化点

  • 计算开销随特征数二次增长,宽表(数百列)上成本高;
  • 对小数据集仍易过拟合,依赖强正则;
  • 每列一个 token,对"列内多值"(如文本型单元格)表达力有限。

3.2 SAINT

2021 年 SAINT 在 intra-sample(特征间)注意力之外,加入inter-sample(样本间)注意力------让一个样本能"参考"同一 batch 内其他样本的表征。

论文原图 (SAINT 注意力机制,来源:arXiv:2106.01342):

图中展示了 SAINT 的注意力结构:先做列间自注意力(intra-sample,特征之间交互),再做行间注意力(inter-sample,样本之间交互),形成层级注意力。

示例代码

python 复制代码
import torch
import torch.nn as nn

class SAINTBlock(nn.Module):
    """一个 SAINT block = 列间注意力 + 行间注意力 + FFN"""
    def __init__(self, embed_dim, n_heads, ffn_dim, dropout=0.1):
        super().__init__()
        # Step 1: 列间自注意力(intra-sample,同 FT-Transformer)
        self.col_attn = nn.MultiheadAttention(embed_dim, n_heads,
                                               batch_first=True)
        self.col_norm = nn.LayerNorm(embed_dim)

        # Step 2: 行间注意力(inter-sample,参考 batch 内其他样本)
        # 把每个样本的 CLS 表征当作一个 token,样本之间做注意力
        self.row_attn = nn.MultiheadAttention(embed_dim, n_heads,
                                               batch_first=True)
        self.row_norm = nn.LayerNorm(embed_dim)

        # FFN
        self.ffn = nn.Sequential(
            nn.Linear(embed_dim, ffn_dim), nn.GELU(),
            nn.Linear(ffn_dim, embed_dim)
        )
        self.ffn_norm = nn.LayerNorm(embed_dim)

    def forward(self, x):
        # x: [B, n_features, embed_dim]
        # --- 列间注意力:特征之间交互 ---
        col_out, _ = self.col_attn(x, x, x)
        x = self.col_norm(x + col_out)              # [B, n_features, embed_dim]

        # --- 行间注意力:样本之间交互 ---
        # 需要交换维度:把 batch 维当作序列维
        # x: [B, n_features, embed_dim] → [n_features, B, embed_dim]
        x_t = x.transpose(0, 1)                     # [n_features, B, embed_dim]
        row_out, _ = self.row_attn(x_t, x_t, x_t)   # 样本间注意力
        x_t = self.row_norm(x_t + row_out)
        x = x_t.transpose(0, 1)                    # 还原 [B, n_features, embed_dim]

        # --- FFN ---
        x = self.ffn_norm(x + self.ffn(x))
        return x

class SAINT(nn.Module):
    def __init__(self, num_continuous, cat_cardinalities,
                 embed_dim=32, n_heads=4, n_layers=2,
                 ffn_dim=64, d_out=1):
        super().__init__()
        # 复用 FT-Transformer 的 Feature Tokenizer
        self.tokenizer = FeatureTokenizer(
            num_continuous, cat_cardinalities, embed_dim
        )
        self.blocks = nn.ModuleList([
            SAINTBlock(embed_dim, n_heads, ffn_dim) for _ in range(n_layers)
        ])
        self.head = nn.Linear(embed_dim, d_out)

    def forward(self, x_num, x_cat):
        tokens = self.tokenizer(x_num, x_cat)       # [B, 1+n, embed_dim]
        for block in self.blocks:
            tokens = block(tokens)
        return self.head(tokens[:, 0])              # [CLS] 输出

# 使用示例:
model = SAINT(
    num_continuous=5, cat_cardinalities=[50, 20, 10],
    embed_dim=32, n_heads=4, n_layers=2, d_out=1
)
# 输入:x_num=[B,5], x_cat=[B,3]
# 关键:行间注意力让当前样本能"看到" batch 内其他样本

举例:在医学诊断中,当前患者症状 token 经过自注意力后,inter-sample 注意力会发现 batch 内有相似症状组合的历史患者,把它们的表征融合进来------相当于一个"软 kNN"嵌在网络里,对噪声鲁棒。

解决了什么问题:之前的架构都是逐样本独立处理,无法利用"相似样本应给相似表征"的归纳偏置。SAINT 把 set transformer 思想引入表格,让样本间信息流动,对噪声/缺失值更鲁棒。

效果:在带噪声、缺失值的鲁棒性基准上表现好,分类性能小幅领先 FT-Transformer;但样本间注意力带来 O(batch²) 复杂度,训练开销显著增大。

优化点

  • 样本间注意力计算成本高,推理时 batch 选择敏感;
  • 相对 FT-Transformer 增益有限,性价比常被质疑;
  • 仍依赖特征 token 化,未解决宽表扩展性。

3.3 阶段评价

这一代的共识是:Transformer 在中等规模数据上能追平甚至偶尔超过 GBDT,但难以稳定压制;更重要的是,复杂的注意力结构未必带来收益。这为下一阶段的"简约反思"埋下伏笔。

下表归纳几个代表架构在典型基准上的相对表现(定性,非精确数字):

架构 vs GBDT 训练成本 可解释性 小数据鲁棒性
MLP 常落后
TabNet 持平~略输
TabTransformer 持平~略输
FT-Transformer 持平~偶尔胜 中高
SAINT 略胜~持平
GBDT (LightGBM) 基准

关键观察 :成本最高、结构最复杂的 SAINT,相对朴素的 GBDT 的优势却很微弱。这让社区开始反思------复杂度与收益是否成正比? 这直接催生了第四阶段。


第四阶段

4.1 ResNet-MLP 强基线

2021 年的论文 "On Embeddings for Numerical Features in Tabular Deep Learning"(Gorishniy et al.)以及随后的综述揭示了一个令人尴尬的事实:一个类似 ResNet 的纯 MLP 架构,只要数值特征嵌入处理得当,在很多基准上能与 FT-Transformer 持平甚至超越。

论文原图 (数值特征嵌入方式对比,来源:arXiv:2106.11959,Appendix):

论文中对比了多种数值特征嵌入方式(无嵌入、周期嵌入、分桶嵌入等),发现嵌入方式比换主干更重要。配以残差连接 + LayerNorm 的 MLP 即可媲美 Transformer。

示例代码

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

class NumericalEmbedding(nn.Module):
    """数值特征嵌入:标量 → 向量(关键组件!)
    论文发现:如何做这一步比换什么主干更重要。
    这里展示 Periodic 周期嵌入 + 分桶嵌入两种方式。"""
    def __init__(self, num_continuous, embed_dim, n_bins=16, mode='periodic'):
        super().__init__()
        self.mode = mode
        if mode == 'periodic':
            # 周期嵌入:用可学习的频率对连续值做正弦变换
            self.freq = nn.Parameter(torch.randn(num_continuous, embed_dim // 2))
            self.phase = nn.Parameter(torch.randn(num_continuous, embed_dim // 2))
        elif mode == 'bucket':
            # 分桶嵌入:把连续值分桶,每桶一个 embedding
            self.bins = nn.Parameter(torch.linspace(-3, 3, n_bins))  # 分位点
            self.bucket_emb = nn.Embedding(n_bins, embed_dim)

    def forward(self, x_num):
        if self.mode == 'periodic':
            # x_num: [B, num_continuous]
            x = x_num.unsqueeze(-1)                   # [B, num, 1]
            freq = self.freq.unsqueeze(0)             # [1, num, embed//2]
            phase = self.phase.unsqueeze(0)
            x_proj = x * freq + phase
            emb = torch.cat([torch.sin(x_proj), torch.cos(x_proj)], dim=-1)
            return emb                                # [B, num, embed_dim]
        elif self.mode == 'bucket':
            # 量化到桶索引
            idx = torch.bucketize(x_num, self.bins)   # [B, num]
            return self.bucket_emb(idx)               # [B, num, embed_dim]

class ResidualBlock(nn.Module):
    """ResNet 风格的残差块"""
    def __init__(self, dim, dropout=0.1):
        super().__init__()
        self.norm = nn.LayerNorm(dim)
        self.linear1 = nn.Linear(dim, dim)
        self.linear2 = nn.Linear(dim, dim)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x):
        h = self.norm(x)
        h = F.relu(self.linear1(h))
        h = self.dropout(h)
        h = self.linear2(h)
        return x + h                                  # 残差连接

class ResNetMLP(nn.Module):
    def __init__(self, num_continuous, cat_cardinalities,
                 embed_dim=32, n_blocks=4, hidden_dim=256, d_out=1):
        super().__init__()
        # 数值嵌入(关键!)
        self.num_emb = NumericalEmbedding(num_continuous, embed_dim, mode='periodic')
        # 类别嵌入
        self.cat_embeddings = nn.ModuleList([
            nn.Embedding(card, embed_dim) for card in cat_cardinalities
        ])
        # 拼接后过 ResNet
        in_dim = (num_continuous + len(cat_cardinalities)) * embed_dim
        self.proj = nn.Linear(in_dim, hidden_dim)
        self.blocks = nn.ModuleList([
            ResidualBlock(hidden_dim) for _ in range(n_blocks)
        ])
        self.norm = nn.LayerNorm(hidden_dim)
        self.head = nn.Linear(hidden_dim, d_out)

    def forward(self, x_num, x_cat):
        num_tokens = self.num_emb(x_num)              # [B, num, embed_dim]
        cat_tokens = [emb(x_cat[:, i]) for i, emb in
                      enumerate(self.cat_embeddings)]
        all_tokens = torch.cat([
            num_tokens.reshape(x_num.shape[0], -1),
            *[ct.unsqueeze(0).expand(x_num.shape[0], -1) for ct in cat_tokens]
        ] if False else
            [num_tokens.reshape(x_num.shape[0], -1)] +
            [emb(x_cat[:, i]) for i, emb in enumerate(self.cat_embeddings)],
            dim=1)
        x = self.proj(all_tokens)
        for block in self.blocks:
            x = block(x)
        return self.head(self.norm(x))

# 使用示例:
model = ResNetMLP(
    num_continuous=5, cat_cardinalities=[50, 20, 10],
    embed_dim=32, n_blocks=4, hidden_dim=256, d_out=1
)
# 关键:性能来自数值嵌入(periodic),不是主干架构!

举例:同样一个银行客户流失数据集,FT-Transformer vs ResNet-MLP 的 AUC 差异可能小于 0.5%,但 ResNet-MLP 训练快数倍、调参更少。很多时候,数值特征是用分桶嵌入(把连续值按分位数分成若干桶,每桶一个 embedding 向量)还是用 Fourier 特征(周期函数编码),比换什么主干更重要。

解决了什么问题 :澄清了"性能来自哪里"。之前社区把 Transformer 的注意力当作制胜法宝,ResNet-MLP 的回归表明:数值嵌入方式 + 训练正则才是关键,而非主干架构的复杂度。这厘清了 NN 的适用边界。

效果:在多数中等规模表格基准上与 FT-Transformer 并驾齐驱,训练成本远低,成为"首选 NN 基线"。

优化点

  • 仍是在目标表上从零训练参数,无法跨表迁移;
  • 对类别嵌入的语义建模不如 Transformer 细腻;
  • 本质是"回归",未提出新范式。

4.2 TabR

2023 年 TabR (Gorishniy et al.)进一步给出结论:在多数表格任务上,一个设计良好的 MLP 系架构就是最佳选择。但 TabR 的核心创新在于检索增强------在 MLP 基础上引入 kNN 式的样本检索。

论文原图 (MLP 基线架构,来源:arXiv:2307.14338):

论文中展示了 MLP 基线的两种形式:无嵌入(连续特征直接标准化)和有嵌入(连续特征经分桶/周期嵌入)。TabR 在此基础上引入了检索模块,从训练集中动态选取最相似的样本参与预测。

示例代码

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

class TabR(nn.Module):
    """TabR: MLP + 检索增强 (retrieval-augmented)
    核心:在 MLP 表征空间做 kNN 检索,把相似样本的信息融入预测。"""
    def __init__(self, num_continuous, cat_cardinalities,
                 embed_dim=16, hidden_dim=256, n_blocks=4,
                 n_neighbors=10, d_out=1):
        super().__init__()
        self.n_neighbors = n_neighbors
        # 特征嵌入(复用 ResNet-MLP 的设计)
        self.num_emb = nn.ModuleList([
            nn.Linear(1, embed_dim) for _ in range(num_continuous)
        ])
        self.cat_embeddings = nn.ModuleList([
            nn.Embedding(card, embed_dim) for card in cat_cardinalities
        ])
        # 编码器 E:把特征序列编码成行向量
        in_dim = (num_continuous + len(cat_cardinalities)) * embed_dim
        self.encoder = nn.Sequential(
            nn.Linear(in_dim, hidden_dim), nn.ReLU(),
            nn.Linear(hidden_dim, hidden_dim), nn.ReLU(),
            nn.Linear(hidden_dim, hidden_dim)
        )
        # 检索 key 的投影
        self.key_proj = nn.Linear(hidden_dim, hidden_dim)
        # 预测器 P:融合目标样本 + 检索结果
        self.predictor = nn.Sequential(
            nn.Linear(hidden_dim * 2, hidden_dim), nn.ReLU(),
            nn.Linear(hidden_dim, d_out)
        )

    def _encode(self, x_num, x_cat):
        """编码特征 → 行向量"""
        num_tokens = [proj(x_num[:, i:i+1]) for i, proj in
                      enumerate(self.num_emb)]       # 每个返回 [B, embed_dim]
        cat_tokens = [emb(x_cat[:, i]) for i, emb in
                      enumerate(self.cat_embeddings)]
        x = torch.cat(num_tokens + cat_tokens, dim=1)  # [B, in_dim]
        return self.encoder(x)                           # [B, hidden_dim]

    def forward(self, x_num, x_cat, train_num, train_cat, train_y):
        """
        x_num/x_cat: 目标样本
        train_num/train_cat/train_y: 训练集(用于检索)
        """
        # 编码目标样本和训练集
        q = self._encode(x_num, x_cat)                  # [B, hidden]
        k = self._encode(train_num, train_cat)          # [N, hidden]
        q_k = self.key_proj(q)                           # [B, hidden]

        # 检索:用目标样本的 key 去匹配训练集
        scores = torch.cdist(q_k, k)                    # [B, N] 距离矩阵
        # 取 top-k 最相似的训练样本
        _, indices = torch.topk(-scores, self.n_neighbors, dim=1)  # [B, k]
        # 聚合检索到的邻居表征
        neighbors = k[indices]                           # [B, k, hidden]
        neighbor_agg = neighbors.mean(dim=1)             # [B, hidden]

        # 融合目标样本表征 + 检索结果
        combined = torch.cat([q, neighbor_agg], dim=1)   # [B, hidden*2]
        return self.predictor(combined)                  # [B, d_out]

# 使用示例:
model = TabR(
    num_continuous=5, cat_cardinalities=[50, 20, 10],
    embed_dim=16, hidden_dim=256, n_neighbors=10, d_out=1
)
# 推理时需要传入训练集用于检索
# pred = model(x_test_num, x_test_cat, x_train_num, x_train_cat, y_train)

举例:在 OpenML-CC18 基准套件的多个数据集上,TabR(无注意力)击败了许多 Transformer 架构,性能却接近 LightGBM。它通过在 MLP 表征空间做 kNN 检索,让模型"参考"训练集中最相似的样本------既保持了 MLP 的简洁,又获得了类似 kNN 的非参数灵活性。

解决了什么问题 :把"简约主义"推到极致,同时引入检索增强机制------在 MLP 表征空间动态检索相似样本,而非用复杂的注意力。它证明了简单的 MLP + 检索比复杂的 Transformer 更有效。

效果:在多数中等规模表格基准上优于或持平于 Transformer 类方法,训练快、调参少,是"务实派"的代表。

优化点

  • 检索依赖训练集,推理时需维护检索索引,工程复杂度高;
  • 仍受限于"在目标表从零训练"的范式,无法跨表迁移;
  • 在需要迁移、零样本的场景(如新业务冷启动)完全失效;
  • 性能天花板由数据量决定,数据不足则无解------这天然引出第五阶段。

第四阶段小结

这一阶段的"反思"并非否定 Transformer,而是厘清了边界:

数据规模 / 特性 更优方向
小到中等样本、异构特征 GBDT 或 ResNet-MLP 仍是首选
大样本、需要迁移/零样本 Transformer / 基础模型开始有价值

核心论断 :在"从零训练参数"的范式下,表格 NN 已触及天花板------无论架构多花哨,小数据上就是打不过 GBDT。要突破,必须换范式:绕开"在目标表上训练参数",转而用先验 + 上下文。这条分水岭自然引出了第五阶段。


第五阶段

5.1 TabPFN

2022 年 Hollmann et al. 提出 TabPFN (Tabular Prior-Data Fitted Network),带来范式革新------把"数据当作参数",零参数更新完成预测

论文原图 (TabPFN 两阶段架构,来源:arXiv:2207.01848):

图中展示了 TabPFN 的两阶段方法:(a) 离线先验拟合阶段------在海量合成数据集上预训练 Transformer,学习"如何解决任意表格分类问题";(b) 在线推理阶段------把训练集样本作为上下文序列输入,权重冻结,直接输出预测。

先验组件(合成数据的分布来源):

图中展示了 TabPFN 先验的构成:(a) 贝叶斯神经网络结构,(b) 结构因果模型示例,© 训练时采样的各种 SCM 结构。先验的多样性决定了模型能处理多广泛的表格分布。

示例代码

python 复制代码
# === TabPFN 的核心思想(概念示意,非可运行代码) ===
# TabPFN 的关键:不在目标数据上训练参数,而是"在上下文中推理"

import torch
import torch.nn as nn

class TabPFNConcept(nn.Module):
    """TabPFN 概念模型:预训练 Transformer 做上下文学习
    实际实现远比这复杂,这里仅展示核心思想。"""
    def __init__(self, embed_dim=256, n_heads=8, n_layers=12):
        super().__init__()
        # 特征编码器:把 (特征值, 标签) 对编码成 token
        self.feature_encoder = nn.Linear(1, embed_dim)  # 简化
        self.label_encoder = nn.Embedding(2, embed_dim) # 二分类标签
        # 预训练的 Transformer(权重冻结!)
        self.transformer = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(
                d_model=embed_dim, nhead=n_heads,
                dim_feedforward=embed_dim * 4,
                batch_first=True, activation='gelu'
            ), n_layers
        )
        self.head = nn.Linear(embed_dim, 2)  # 分类输出

    def forward(self, train_features, train_labels, test_features):
        """
        关键:不更新任何参数!
        train_features: [N_train, n_feat] --- 训练集
        train_labels: [N_train] --- 训练标签
        test_features: [N_test, n_feat] --- 测试集
        """
        # --- 构造上下文序列 ---
        # 把训练集每行的 (特征, 标签) 编码成 token
        context_tokens = []
        for i in range(train_features.shape[0]):
            feat_tokens = [self.feature_encoder(f) for f in
                           train_features[i]]  # 每列一个 token
            label_token = self.label_encoder(train_labels[i])
            context_tokens += feat_tokens + [label_token]

        # 测试样本:只编码特征(没有标签!)
        test_tokens = []
        for i in range(test_features.shape[0]):
            feat_tokens = [self.feature_encoder(f) for f in
                          test_features[i]]
            test_tokens += feat_tokens

        # 拼接成完整序列
        all_tokens = torch.stack(context_tokens + test_tokens,
                                 dim=0).unsqueeze(0)  # [1, seq_len, embed_dim]

        # --- 过 Transformer(权重冻结,不反向传播!)---
        with torch.no_grad():  # 关键:零参数更新
            output = self.transformer(all_tokens)

        # 取测试样本对应的输出位置做预测
        return output  # 概念示意

# === 实际使用(pip install tabpfn)===
# from tabpfn import TabPFNClassifier
# clf = TabPFNClassifier()           # 加载预训练权重
# clf.fit(X_train, y_train)          # 仅构造上下文,不训练参数
# pred = clf.predict(X_test)         # 零样本预测,秒级出结果

举例 :给 TabPFN 喂 100 条带标签的客户数据 + 一条新客户特征,它像 GPT 接龙一样直接输出违约概率------零参数更新。在 OpenML 的小数据分类基准上,TabPFN 常超越 XGBoost 和所有微调式 NN。其本质是"先验数据即学习参数":把"什么样的数据分布和任务结构是常见的"这件事,压缩进预训练权重里。

解决了什么问题:突破了第四阶段的范式天花板。之前所有架构都在"目标表上从零学参数",小数据必然过拟合;TabPFN 把对"表格任务本身"的先验知识预训练进权重,推理时靠上下文学习(in-context learning, ICL)直接迁移,绕开了过拟合困境。

效果 :在小数据分类(≤ 1000 样本、≤ 100 特征)上常超越 XGBoost 和所有微调 NN,且无需调参、无需训练,几秒内出结果。被视为表格领域的"基础模型"雏形。

优化点

  • 上下文长度受限,通常 ≤ 10k 样本,中等以上数据集无法直接用;
  • 限于分类,回归能力相对弱;
  • 先验来自合成数据,对分布外(OOD)的真实工业数据泛化存疑;
  • 仅处理"扁平表格",不感知多表关系/时间序列。

5.2 TabICL 与检索增强

TabPFN 的局限在于上下文长度------Transformer 的注意力复杂度随序列长度二次增长,通常只能处理 ≤ 10k 样本的上下文。后续工作沿两条互补的技术路线扩展:

路线 核心思路 代表工作
路线 A:扩大上下文容量 设计更强的行编码器,把整行压缩成一个固定维度向量,再用 Transformer 处理"行序列"而非"特征序列",从而在同等显存下塞入更多样本 TabICL(Qu et al., 2025)
路线 B:检索增强 不把所有训练样本塞进上下文,而是先用检索器选出与当前预测样本最相关的 k 个,只把这 k 个送入 ICL 模型 检索增强 ICL(NeurIPS 2024)

路线 A 的关键设计------分布感知列嵌入 (TabICL,来源:arXiv:2502.05564):

怎么读这张图 :TabICL 对每个特征列生成两种嵌入并拼接

  • 值嵌入(Value Embedding):编码该列在当前样本的具体取值(如"年龄=35"),类似于 FT-Transformer 的 token 化
  • 分布嵌入(Distribution Embedding) :编码该列的全局统计特性------均值、方差、分位数等。这让模型不仅知道"这个值是多少",还知道"这一列整体长什么样",相当于给每列一个"列档案"

为什么这样设计 :TabPFN 把每个 (特征, 值) 当作独立 token,序列长度 = 样本数 × 列数,10k 样本 × 100 列 = 100 万 token,Transformer 无法承受。TabICL 的列嵌入把一行压缩成一个固定维度向量(而非每列一个 token),序列长度降到 = 样本数,同等显存下能处理 60k 样本(TabPFN 的 6 倍)。

路线 A 的效果------TabICL 相对 MLP 的性能提升

怎么读这张图:横轴是数据集的样本量,纵轴是 TabICL 相对于 MLP 基线的相对性能提升(>0 表示 TabICL 更好)。

  • 小数据(< 1k 样本):提升有限,因为上下文样本太少,ICL 优势发挥不出来
  • 中等数据(1k-60k 样本):提升最明显,这正是 ICL 的"甜区"------有足够上下文样本支撑推理,又没超出上下文容量
  • 大数据(> 60k):曲线回落,因为开始触及上下文容量上限,无法纳入全部训练样本

与 TabPFN 的差异 :TabPFN 在 < 1k 样本时最强(上下文小、推理快),但 > 10k 就力不从心;TabICL 通过列嵌入压缩把上限推到 60k,覆盖了 TabPFN 够不到的中等数据区间。两者是互补关系而非替代。

路线 B 的示意代码(检索增强上下文模型):

python 复制代码
import torch
import torch.nn as nn

class RetrievalAugmentedICL(nn.Module):
    """检索增强的上下文学习模型(概念示意)
    当训练集太大塞不进上下文时,先检索最相关的 k 个样本。"""
    def __init__(self, embed_dim=256, n_heads=8, n_layers=6,
                 n_retrieved=128):
        super().__init__()
        self.n_retrieved = n_retrieved
        # 行编码器:把一行的特征编码成一个向量
        self.row_encoder = nn.Sequential(
            nn.Linear(1, embed_dim)  # 简化:实际用列嵌入
        )
        # 检索器:学习型 kNN
        self.retrieval_proj = nn.Linear(embed_dim, embed_dim)
        # 上下文 Transformer
        self.transformer = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(
                d_model=embed_dim, nhead=n_heads,
                dim_feedforward=embed_dim * 4,
                batch_first=True, activation='gelu'
            ), n_layers
        )
        self.head = nn.Linear(embed_dim, 2)

    def forward(self, query_row, train_rows, train_labels):
        """
        query_row: [1, n_feat] --- 目标样本
        train_rows: [N, n_feat] --- 训练集(可能很大)
        train_labels: [N] --- 训练标签
        """
        # --- 阶段1: 编码 ---
        q = self.row_encoder(query_row).mean(dim=1)    # [1, embed_dim]
        k = self.row_encoder(train_rows).mean(dim=1)  # [N, embed_dim]

        # --- 阶段2: 检索 top-k 最相关样本 ---
        scores = torch.matmul(
            self.retrieval_proj(q), k.T                 # [1, N]
        ) / (embed_dim ** 0.5)
        _, topk_idx = torch.topk(scores, self.n_retrieved, dim=1)  # [1, k]
        retrieved = k[topk_idx.squeeze(0)]               # [k, embed_dim]
        retrieved_labels = train_labels[topk_idx.squeeze(0)]

        # --- 阶段3: 构造上下文序列 ---
        context = torch.cat([retrieved, q], dim=0).unsqueeze(0)  # [1, k+1, embed_dim]
        out = self.transformer(context)
        return self.head(out[:, -1])                     # 取最后位置(query)输出

举例:一个 50 万行的电商用户分群任务,远超 TabPFN 的 10k 上下文。

  • 路线 A(TabICL):用列嵌入把每行压缩成固定维度向量,序列长度从 500 万 token 降到 50 万,配合更强的预训练直接处理------但仍可能在 50 万行时触及上限
  • 路线 B(检索增强):先用检索器从 50 万行中选出与当前预测样本最相关的 1000 行,只把这 1000 行送入 ICL 模型------绕开了上下文容量限制,但检索质量决定上限

两条路线可组合使用:先用检索缩减候选集,再用 TabICL 的高容量编码器处理检索结果。

解决了什么问题:把上下文学习范式从"小数据专用"扩展到中等乃至较大数据,打破了 TabPFN 的规模天花板。

效果:在中等规模数据上开始展现竞争力;检索增强路线在 NeurIPS 2024 有相关工作,让 ICL 模型逼近甚至持平微调基线。

优化点

  • 检索器的质量决定上限,本身又是一个需训练/调优的组件;
  • 长上下文 Transformer 的计算/显存开销大;
  • 仍是"扁平表格"假设,对多表/时序数据无结构感知。

5.3 TabFM 与商业基础模型

Google 等推出的 TabFM(Tabular Foundation Model)代表了"工业级表格基础模型"的尝试。

TabFM 概念图 (TabFM 的范式定位,来源:joseparreogarcia.substack.com):

图中展示了 TabFM 如何融合 TabPFN、TabICL 和 TabDPT 的思想:行压缩到固定维度向量,再用 ICL Transformer 读取整表作为上下文。

TabFM 输入示例来源同上):

图中展示了 TabFM 的输入格式:整张表(训练行 + 查询行)作为一个上下文送入模型,模型直接输出预测。

示例代码(使用 Google 官方 TabFM 库):

python 复制代码
# === 实际使用 TabFM(pip install tabfm)===
import numpy as np
import pandas as pd
from tabfm import TabFMClassifier
from tabfm import tabfm_v1_0_0_pytorch as tabfm_v1_0_0

# 加载预训练权重(从 HuggingFace 自动下载)
model = tabfm_v1_0_0.load()

# 初始化 scikit-learn 兼容的分类器
clf = TabFMClassifier(model=model)

# 准备数据(支持混合连续和类别特征)
X_train = pd.DataFrame({
    "age": [25.0, 45.0, 35.0, 50.0],
    "job": ["engineer", "manager", "engineer", "manager"],
    "income": [80000, 120000, 90000, 130000]
})
y_train = np.array(["low_risk", "high_risk", "low_risk", "high_risk"])

X_test = pd.DataFrame({
    "age": [30.0, 48.0],
    "job": ["engineer", "manager"],
    "income": [85000, 125000]
})

# Fit:仅准备编码器(序数编码 + 数值标准化),不训练模型参数!
clf.fit(X_train, y_train)

# Predict:零样本预测------模型读取训练集作为上下文
predictions = clf.predict(X_test)
# 输出: array(['low_risk', 'high_risk'])

# 关键:整个过程中模型权重没有任何更新!
# 模型通过 in-context learning 从训练集"理解"任务后直接预测。

举例:一个新业务线只有几百条标注样本,传统做法要么用 GBDT 过拟合,要么人工标更多数据。用 TabFM 可以零样本/少样本直接预测,或把它当特征提取器,输出表 embedding 喂给下游轻量模型。

解决了什么问题:回答"是否值得做通用表格预训练"。TabNet、TabTransformer、SAINT、FT-Transformer 证明了 NN 在表格上"可行",但都是单表从零训练;TabFM 试图沉淀跨表通用先验,提供开箱即用、可零样本/少样本迁移的工业级模型。

效果:作为 2024-2025 的新方向,已在部分内部基准展现零样本迁移能力,但通用性与成熟度仍在验证中。

优化点

  • 预训练数据的分布覆盖度决定泛化上限;
  • 评测标准未统一,跨表迁移的"真实"提升仍有争议;
  • 与 LLM 路线的边界尚未厘清(见 5.4)。

5.4 LLM 与表格理解

另一条线是大语言模型直接做表格任务。

架构说明(LLM 表格理解范式):

复制代码
原始表格 + 自然语言问题
    → LLM (TableGPT2 / 通用 LLM)
    → 理解表格结构 + 语义
    → 自然语言问数 / 数据摘要 / 可视化 / 报表解读

示例代码(使用 LLM 做表格问答):

python 复制代码
# === 使用 LLM 做表格问答(概念示意)===

import pandas as pd

# 模拟一张销售明细表
sales_df = pd.DataFrame({
    "区域": ["华东", "华东", "华南", "华南", "华北", "华北"],
    "产品": ["A类", "B类", "A类", "B类", "A类", "B类"],
    "季度": ["Q1", "Q1", "Q1", "Q1", "Q1", "Q1"],
    "销售额": [120, 85, 200, 150, 90, 110],
    "环比变化": [0.05, -0.12, 0.15, -0.08, 0.02, -0.20]
})

# 构造 prompt
question = "华东区上季度哪类产品环比降幅最大?"

prompt = f"""
你是一个数据分析助手。请根据以下表格数据回答问题。

表格数据(CSV 格式):
{sales_df.to_csv(index=False)}

问题:{question}

请给出答案和计算过程。
"""

# 调用 LLM(这里用概念示意,实际用 OpenAI/通义千问等 API)
# response = llm.chat(prompt)
# 预期回答:
# "华东区 Q1 中 B 类产品环比降幅最大,为 -12%(从 96.6 降至 85)。
#  A 类产品环比 +5%(从 114.3 升至 120)。"

# === 关键区别 ===
# TFM 路线 (TabPFN/TabFM): y = f(x; D),学习从特征到标签的映射
# LLM 路线 (TableGPT2): 自然语言 → 表格理解 → 自然语言回答
# 两者的融合方向:既懂语义又能精准预测

举例:用户上传一张销售明细表,问"华东区上季度哪类产品环比降幅最大?" TableGPT2 直接读表结构、算环比、给出答案,无需用户写 SQL。

解决了什么问题:降低表格分析的交互门槛------用自然语言直接问数、解读报表,而非写 SQL/调模型。

效果:在表格问答、报表解读类任务上进展快,2024-2026 涌现 TableGPT2 等模型。

优化点

  • LLM 路线偏理解与交互 (问数、解读),TFM 路线偏预测建模(y=f(x;D));
  • 长表/大表的 token 化成本高,精度受上下文限制;
  • 两条路线正在融合------未来可能出现"既懂语义又能精准预测"的统一表格基础模型。

5.5 三大表格基础模型对比:TabPFN vs TabICL vs TabFM

三者同属"表格基础模型"路线,但定位、技术路线和适用边界差异显著。

核心差异一览
维度 TabPFN (2022) TabICL (2025) TabFM (2025)
提出方 Hollmann et al.(学术) Qu et al.(学术) Google Research(工业)
定位 小数据 ICL 的概念验证 中等数据 ICL 的扩展 工业级通用表格基础模型
上下文容量 ≤ 10k 样本 ≤ 60k 样本 待确认(官方称"可读整表")
预训练数据 合成数据(SCM/BNN 先验) 合成数据(更大规模) 大规模真实表格数据
行编码 逐 (特征,值) 做 token 分布感知列嵌入,行压缩成单向量 行压缩 + 列感知编码
任务类型 分类为主,回归较弱 分类+回归 分类+回归
推理速度 秒级(小上下文) 秒~十秒级 秒级
开源状态 开源(HuggingFace) 开源(GitHub) 开源权重(非商用许可)
技术路线差异

#mermaid-svg-E9XcNL47oC2hGZHW{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-E9XcNL47oC2hGZHW .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-E9XcNL47oC2hGZHW .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-E9XcNL47oC2hGZHW .error-icon{fill:#552222;}#mermaid-svg-E9XcNL47oC2hGZHW .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-E9XcNL47oC2hGZHW .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-E9XcNL47oC2hGZHW .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-E9XcNL47oC2hGZHW .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-E9XcNL47oC2hGZHW .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-E9XcNL47oC2hGZHW .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-E9XcNL47oC2hGZHW .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-E9XcNL47oC2hGZHW .marker{fill:#333333;stroke:#333333;}#mermaid-svg-E9XcNL47oC2hGZHW .marker.cross{stroke:#333333;}#mermaid-svg-E9XcNL47oC2hGZHW svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-E9XcNL47oC2hGZHW p{margin:0;}#mermaid-svg-E9XcNL47oC2hGZHW .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-E9XcNL47oC2hGZHW .cluster-label text{fill:#333;}#mermaid-svg-E9XcNL47oC2hGZHW .cluster-label span{color:#333;}#mermaid-svg-E9XcNL47oC2hGZHW .cluster-label span p{background-color:transparent;}#mermaid-svg-E9XcNL47oC2hGZHW .label text,#mermaid-svg-E9XcNL47oC2hGZHW span{fill:#333;color:#333;}#mermaid-svg-E9XcNL47oC2hGZHW .node rect,#mermaid-svg-E9XcNL47oC2hGZHW .node circle,#mermaid-svg-E9XcNL47oC2hGZHW .node ellipse,#mermaid-svg-E9XcNL47oC2hGZHW .node polygon,#mermaid-svg-E9XcNL47oC2hGZHW .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-E9XcNL47oC2hGZHW .rough-node .label text,#mermaid-svg-E9XcNL47oC2hGZHW .node .label text,#mermaid-svg-E9XcNL47oC2hGZHW .image-shape .label,#mermaid-svg-E9XcNL47oC2hGZHW .icon-shape .label{text-anchor:middle;}#mermaid-svg-E9XcNL47oC2hGZHW .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-E9XcNL47oC2hGZHW .rough-node .label,#mermaid-svg-E9XcNL47oC2hGZHW .node .label,#mermaid-svg-E9XcNL47oC2hGZHW .image-shape .label,#mermaid-svg-E9XcNL47oC2hGZHW .icon-shape .label{text-align:center;}#mermaid-svg-E9XcNL47oC2hGZHW .node.clickable{cursor:pointer;}#mermaid-svg-E9XcNL47oC2hGZHW .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-E9XcNL47oC2hGZHW .arrowheadPath{fill:#333333;}#mermaid-svg-E9XcNL47oC2hGZHW .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-E9XcNL47oC2hGZHW .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-E9XcNL47oC2hGZHW .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-E9XcNL47oC2hGZHW .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-E9XcNL47oC2hGZHW .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-E9XcNL47oC2hGZHW .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-E9XcNL47oC2hGZHW .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-E9XcNL47oC2hGZHW .cluster text{fill:#333;}#mermaid-svg-E9XcNL47oC2hGZHW .cluster span{color:#333;}#mermaid-svg-E9XcNL47oC2hGZHW div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-E9XcNL47oC2hGZHW .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-E9XcNL47oC2hGZHW rect.text{fill:none;stroke-width:0;}#mermaid-svg-E9XcNL47oC2hGZHW .icon-shape,#mermaid-svg-E9XcNL47oC2hGZHW .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-E9XcNL47oC2hGZHW .icon-shape p,#mermaid-svg-E9XcNL47oC2hGZHW .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-E9XcNL47oC2hGZHW .icon-shape .label rect,#mermaid-svg-E9XcNL47oC2hGZHW .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-E9XcNL47oC2hGZHW .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-E9XcNL47oC2hGZHW .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-E9XcNL47oC2hGZHW :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 扩大容量
真实数据 + 工业化
TabFM: 工业级预训练
训练样本 + 查询行
行压缩 + 列感知编码
整表作为上下文
ICL Transformer

(真实数据预训练)
TabICL: 行级压缩
训练样本

(x₁,y₁)...(xₖ,yₖ)
值嵌入 + 分布嵌入 → 拼接
整行压缩成 1 个向量

序列长度 = k
Transformer

(上下文上限 ~60k样本)
TabPFN: 逐特征 token 化
训练样本

(x₁,y₁)...(xₖ,yₖ)
每个 (特征,值) → 一个 token
序列长度 = k × n_features

10k样本×100列 = 100万token
Transformer

(上下文上限 ~10k样本)

各自优势

TabPFN 的优势

  • 小数据王者:在 ≤ 1k 样本的分类任务上常超越 XGBoost 和所有微调 NN
  • 零调参零训练:加载预训练权重 → 喂数据 → 出结果,几秒完成
  • 先验设计优雅:用 SCM(结构因果模型)生成合成训练数据,让模型学到因果关系而非表面相关

TabICL 的优势

  • 中等数据覆盖:通过行级压缩把容量从 10k 推到 60k,覆盖了 TabPFN 够不到的区间
  • 分布感知嵌入:不仅编码"这个值是多少",还编码"这一列整体分布长什么样",列语义更丰富
  • 与 TabPFN 互补:小数据用 TabPFN,中等数据用 TabICL,两者接力

TabFM 的优势

  • 真实数据预训练:用大规模真实表格数据(而非纯合成),对工业分布的泛化更可信
  • 开箱即用 :scikit-learn 兼容 API,fit → predict 三行代码,工程友好
  • Google 工程化背书:有持续的权重更新和生态支持预期
各自局限

TabPFN 的局限

  • 上下文容量硬上限 ~10k,中等以上数据集无法直接用
  • 先验来自合成数据,对 OOD(分布外)真实工业数据泛化存疑
  • 回归能力明显弱于分类

TabICL 的局限

  • 60k 仍是上限,大数据集需配合检索增强
  • 仍是合成数据预训练,真实分布覆盖有限
  • 学术项目,工程成熟度不及 TabFM

TabFM 的局限

  • 2025 年新发布,跨表迁移的"真实"提升仍有争议,评测标准未统一
  • 预训练权重非商用许可,商用受限
  • 与 TabPFN/TabICL 的性能对比尚缺大规模独立基准
选型决策树
你的场景 推荐
≤ 1k 样本分类,要快 TabPFN
1k-60k 样本,分类或回归 TabICL
需要工业级稳定 + 商用许可可控 TabFM(确认许可后)
> 60k 样本 检索增强 + 上述任一,或回退 GBDT
需要跨表迁移做基准 三者都试,取最优

一句话 :TabPFN 验证了"表格 ICL 可行",TabICL 把它扩展到中等数据,TabFM 把它工业化。三者是递进互补关系,而非替代------选哪个取决于你的数据规模和工程需求。


全景总结

回看整条脉络,是一个螺旋上升的过程:
#mermaid-svg-oWv4FnWuB5DPUJjz{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-oWv4FnWuB5DPUJjz .error-icon{fill:#552222;}#mermaid-svg-oWv4FnWuB5DPUJjz .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-oWv4FnWuB5DPUJjz .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-oWv4FnWuB5DPUJjz .marker{fill:#333333;stroke:#333333;}#mermaid-svg-oWv4FnWuB5DPUJjz .marker.cross{stroke:#333333;}#mermaid-svg-oWv4FnWuB5DPUJjz svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-oWv4FnWuB5DPUJjz p{margin:0;}#mermaid-svg-oWv4FnWuB5DPUJjz .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-oWv4FnWuB5DPUJjz .cluster-label text{fill:#333;}#mermaid-svg-oWv4FnWuB5DPUJjz .cluster-label span{color:#333;}#mermaid-svg-oWv4FnWuB5DPUJjz .cluster-label span p{background-color:transparent;}#mermaid-svg-oWv4FnWuB5DPUJjz .label text,#mermaid-svg-oWv4FnWuB5DPUJjz span{fill:#333;color:#333;}#mermaid-svg-oWv4FnWuB5DPUJjz .node rect,#mermaid-svg-oWv4FnWuB5DPUJjz .node circle,#mermaid-svg-oWv4FnWuB5DPUJjz .node ellipse,#mermaid-svg-oWv4FnWuB5DPUJjz .node polygon,#mermaid-svg-oWv4FnWuB5DPUJjz .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-oWv4FnWuB5DPUJjz .rough-node .label text,#mermaid-svg-oWv4FnWuB5DPUJjz .node .label text,#mermaid-svg-oWv4FnWuB5DPUJjz .image-shape .label,#mermaid-svg-oWv4FnWuB5DPUJjz .icon-shape .label{text-anchor:middle;}#mermaid-svg-oWv4FnWuB5DPUJjz .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-oWv4FnWuB5DPUJjz .rough-node .label,#mermaid-svg-oWv4FnWuB5DPUJjz .node .label,#mermaid-svg-oWv4FnWuB5DPUJjz .image-shape .label,#mermaid-svg-oWv4FnWuB5DPUJjz .icon-shape .label{text-align:center;}#mermaid-svg-oWv4FnWuB5DPUJjz .node.clickable{cursor:pointer;}#mermaid-svg-oWv4FnWuB5DPUJjz .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-oWv4FnWuB5DPUJjz .arrowheadPath{fill:#333333;}#mermaid-svg-oWv4FnWuB5DPUJjz .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-oWv4FnWuB5DPUJjz .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-oWv4FnWuB5DPUJjz .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-oWv4FnWuB5DPUJjz .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-oWv4FnWuB5DPUJjz .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-oWv4FnWuB5DPUJjz .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-oWv4FnWuB5DPUJjz .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-oWv4FnWuB5DPUJjz .cluster text{fill:#333;}#mermaid-svg-oWv4FnWuB5DPUJjz .cluster span{color:#333;}#mermaid-svg-oWv4FnWuB5DPUJjz div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-oWv4FnWuB5DPUJjz .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-oWv4FnWuB5DPUJjz rect.text{fill:none;stroke-width:0;}#mermaid-svg-oWv4FnWuB5DPUJjz .icon-shape,#mermaid-svg-oWv4FnWuB5DPUJjz .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-oWv4FnWuB5DPUJjz .icon-shape p,#mermaid-svg-oWv4FnWuB5DPUJjz .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-oWv4FnWuB5DPUJjz .icon-shape .label rect,#mermaid-svg-oWv4FnWuB5DPUJjz .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-oWv4FnWuB5DPUJjz .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-oWv4FnWuB5DPUJjz .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-oWv4FnWuB5DPUJjz :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 复杂未必赢
换范式: 先验+上下文
阶段一

MLP / Wide & Deep / DeepFM

解决: 特征怎么进网络
阶段二

TabNet / TabTransformer

引入注意力 / 类别语义化
阶段三

FT-Transformer / SAINT

统一 token 化 / 样本间注意力
阶段四

ResNet-MLP / TabR

简约反思: 厘清参数学习天花板
阶段五

TabPFN / TabFM / TableGPT2

上下文学习 / 表格基础模型

阶段 时间 代表架构 核心思想 对 GBDT
~2016 MLP、Wide & Deep、DeepFM 显式/隐式特征交互分工 推荐场景已可一战
2019-2020 TabNet、TabTransformer 引入注意力 / 类别嵌入语义化 仍未稳定超越
2021 FT-Transformer、SAINT 统一 token 化、样本间注意力 中等规模追平
2021-2023 ResNet-MLP、TabR 简约反思、数值嵌入才是关键 厘清了 NN 的适用边界
2022-2026 TabPFN、TabICL、TabFM、TableGPT2 上下文学习、表格基础模型 小数据零样本超越;大模型路线兴起

演进逻辑

  • 第一、二阶段在解决"特征怎么进网络";
  • 第三阶段想用 Transformer 统一一切,却发现复杂未必赢;
  • 第四阶段的反思厘清了参数学习路线的天花板------它只在数据足够、需要迁移时才真正占优;
  • 第五阶段绕开"在目标表上训练参数"的局限,用先验 + 上下文打开新空间。

给从业者的选型建议

对从业者而言,务实的判断至今仍是:

你的场景 推荐方案
小到中等规模异构表格(< 10 万样本) LightGBM/XGBoost 为首选;若要 NN 强基线,用 ResNet-MLP / TabR
中等规模,追求精度极限 FT-Transformer + 强正则,可与 GBDT 做集成
小数据分类,无标注预算 TabPFN 零样本试一把
需要跨表迁移 / 零样本 TabFM / TabPFN 路线(仍在成熟中)
需要自然语言问数 / 报表解读 TableGPT2 / LLM 路线
亿级样本推荐 / CTR DeepFM 及其变体仍是工业事实标准

一句话:架构的选择,终究由数据的规模与分布决定,而非由架构的时髦程度决定。


深度问答:两个核心疑问的澄清

以下整理两个常见疑问,帮助读者更深入理解各架构的设计动机与局限。

Q1:为什么 MLP 无法自动发现高阶特征交互?加深网络行不行?

什么是高阶交互? 一个特征对预测的贡献依赖于其他特征的取值,就是交互。参与耦合的特征越多,"阶"越高。

  • 一阶(主效应)面积 越大房价越高,与地段无关 → w₁ × 面积
  • 二阶交互面积 每平米值多少钱取决于地段 → 需要 面积 × 地段 乘法项
  • 三阶交互男 AND 20-30 AND 游戏 三个特征同时取特定值才有强信号

为什么线性层不行? 线性层 y = Σ wᵢ·xᵢ 是纯加法,x₁ 的贡献 w₁·x₁ 永远与 x₂ 无关。它只能表达"每个特征独立贡献",无法表达"一个特征的贡献随另一个特征变化"。要表达交互,必须有乘法项(x₁×x₂),而线性层里没有。

加深 MLP 理论上能逼近交互吗? 能。ReLU 是通用函数逼近器,x₁×x₂ 可以用一组 ReLU 折线逼近。但实践上不可靠,原因有三:

困难 说明
归纳偏置不匹配 MLP 的结构假设是"加法 → 非线性 → 加法",没有任何组件提示"去找乘法关系"。交互只能通过非线性复合间接涌现,梯度信号极弱
表格数据小,深网过拟合 图像能用 50 层 ResNet 是因为有空间局部性可复用参数;表格无结构可利用,加深 → 参数爆炸 → 过拟合,3-4 层后收益递减甚至变负
类别交互的稀疏梯度 one-hot 输入大部分为 0。某个特定三元组(如"男+20-30+游戏")只有 < 5% 样本满足条件,对应神经元绝大多数训练步里梯度为 0,学不动

这就是 Wide 部分和 FM 存在的原因

  • Wide 部分 :把 AND(男, 20-30, 游戏) 做成一个手工二值特征 φ,直接给线性层一个权重 w_φ → 交互被显式提供,不需"发现"
  • FM :每个特征有 embedding 向量 e,二阶交互 = e₁·e₂(内积)→ 自动参数化所有两两交互,但仅二阶

结论 :MLP 不是"不能"表达交互,而是"不擅长"。各种专用架构(Wide 手工交叉、FM 内积、Attention 加权)都是在给网络注入正确的交互先验,让模型直接参数化交互而非从零发现。


Q2:为什么 FT-Transformer 不能像 LLM 一样训一次、处处迁移?

结论先行 :FT-Transformer 必须每个场景从头训练,无法实现 LLM 式的通用迁移。根源是表格数据没有跨数据集共享的离散符号空间。

核心对比

LLM FT-Transformer
基本单元 词/子词("猫""king") 列 + 列值的组合
符号跨任务一致? ✅ "猫"在任何文本里都指猫 ❌ 数据集 A 的"收入=50000"与数据集 B 的"收入=50000"语义不同
词表大小 固定(~50K-100K) 每个数据集不同(列数 × 各列基数)
Embedding 可迁移? ✅ 预训练语义直接复用 ❌ 只在当前数据集有意义

为什么 LLM 能迁移? 自然语言有全局共享的离散符号系统。"猫"这个 token 在维基、新闻、小说里都指同一个东西,预训练学到的 embedding 在任何任务里都能复用。

为什么 FT-Transformer 不能迁移? 表格的"符号"是列名 × 列值的笛卡尔积,每个数据集都不同:

  • 数据集 A(银行风控)有列 职业(50 种值)→ Embedding 表 50 行
  • 数据集 B(电商 CTR)有列 城市(30 种值)→ Embedding 表 30 行
  • 两个 Embedding 表完全不可共享------列名不同、值不同、语义不同

连续特征更无"词表"可言:FT-Transformer 用 Linear(1, d) 把标量投影成向量,这个权重是特定数据集特定列的投影矩阵,换数据集即失效。

这正是第四阶段→第五阶段的突破点

范式 如何处理新数据集 迁移能力
FT-Transformer(参数学习) 重头训练所有参数 ❌ 无法迁移
TabPFN/TabFM(上下文学习) 不训参数,换上下文即可 ✅ 权重不变,处处可用

类比:FT-Transformer 像"背答案的学生"------把知识记进参数,换题就不会了;TabPFN 像"会看参考题的学生"------不背答案,但学会了"看几道例题就能推理"的元能力。


参考来源

本文基于以下公开资料整理,时间截至 2026 年 8 月。所有架构图均来自原始论文或官方资料:


本文旨在梳理技术脉络,架构之间的性能比较随数据集、超参、预处理而变,具体选型请以自有数据实测为准。所有架构图版权归原作者所有,仅作技术说明引用。

相关推荐
MacroZheng1 小时前
面试官皱眉:"你懂 Vibe Coding,那你说superpowers和grill-me怎么选?",我:"小孩才做选择,我全都要!"
java·人工智能·后端
OceanBase数据库官方博客1 小时前
深度拆解seekdb:AI Native Database 的技术架构与核心能力
数据库·人工智能·架构
大模型码小白1 小时前
【AI大模型】DeepSeek Harness 深度解析:大模型评测框架的架构与实践
java·运维·人工智能·spring·架构·自动化
码农小旋风1 小时前
GPT-6 Astra 直接进 Blender,AI 开始从会说变成会做
人工智能·gpt·blender
武子康1 小时前
9/10 和 900/1000 都是 90%:机器人成功率计算器必须阻止的误判
人工智能·llm·agent
weixin_618232301 小时前
OpenAI智能体“劫持”德国老网站当留言板
人工智能·安全
未来之窗软件服务1 小时前
城市街道社区综合管理开发之档案管理—东方仙盟
人工智能·仙盟创梦ide·档案系统
武哥聊编程1 小时前
【AI实战项目】基于SpringAI+Springboot+Vue的AI智能简历优化助手
vue.js·人工智能·spring boot
大虾别跑1 小时前
ai-daily-2026-09-07
人工智能