南非世界杯时间

【表格数据建模】Mambular: 一种用于表格深度学习的序列模型

Mambular: 一种用于表格深度学习的序列模型

https://arxiv.org/pdf/2408.06291

摘要

表格数据的分析传统上由梯度增强决策树(GBDTs)主导,因其在处理混合类别和数值特征方面的高效性而闻名。然而,最近的深度学习创新正在挑战这种主导地位。我们介绍Mambular,这是针对表格数据优化的Mamba架构的一个适应版本。我们对Mambular与包括神经网络和基于树的方法在内的最先进模型进行了广泛基准测试,并展示了其在各种数据集中的竞争性能。此外,我们探索了Mambular的不同适应性,以理解其对表格数据的有效性。我们研究了不同的池化策略、特征交互机制和双向处理。我们的分析表明,将特征解释为序列并通过Mamba层传递,结果产生了表现出人意料的强劲模型。结果突显了Mambular作为一种多功能和强大的表格数据分析架构的潜力,扩展了深度学习在该领域应用的范围。源代码可在 https://github.com/basf/mamba-tabular 获取。

1 引言

梯度增强决策树(GBDTs)历史上在表格数据分析领域占据主导地位,广泛使用的变体包括XGBoost、LightGBM和CatBoost。这些模型在处理表格数据典型的类别和数值特征混合方面表现出色 (Grinsztajn et al., 2022)。历史上,由于表格数据的内在复杂性和多样性,深度学习模型在处理表格数据时往往表现不佳,通常无法超越GBDTs,因为缺乏处理缺失值、各种特征类型的能力以及进行大量预处理的必要性 (Borisov et al. 2022)。然而,最近在深度学习方面的进展逐渐挑战了这一范式,通过引入创新架构来利用先进机制捕捉复杂特征依赖性, promise significant improvements (Popov et al. 2019; Hollmann et al., 2022; Gorishniy et al., 2021)。

在表格深度学习中的一个最有效的进展是应用注意力机制于像TabTransformer (Huang et al. 2020), FT-Transformer (Gorishniy et al. 2021)和许多其他模型 (Wang and Sun, 2022, Thielmann et al., 2024b; Arik and Pfister, 2021)。这些模型利用注意力机制捕捉特征之间的依赖关系,提供了对传统方法的显著改善。FT-Transformer在各种表格数据集上表现出强劲的性能,通常超越了GBDTs的准确性 (McElfresh et al. 2024)。

此外,更传统的模型如多层感知器(MLP)和残差网络(ResNets)在设计良好并且数据经过彻底预处理时表现出改进 (Gorishniy et al., 2021, 2022)。这些模型特别受益于先进预处理方法的创新,使它们更具竞争力。

最近,Mamba架构 (Gu and Dao, 2023) 在文本问题上显示出良好的结果。此前由Transformer架构主导的任务,如DNA建模和大型语言模型(LLMs),在应用Mamba模型后也得到了改进 (Gu and Dao, 2023; Schiff et al., 2024, Zhao et al., 2024)。

几种适应展示了其多功能性,例如用于图像分类的Vision Mamba (Xu et al. 2024)、视频分析 (Yang et al. 2024; Yue and Li, 2024) 和点云分析 (Zhang et al. 2024, Liu et al. 2024)。此外,该架构还被适应用于时间序列问题,Patro和Agneeswaran (2024)、Wang等 (2024) 和Ahamed与Cheng (2024b)也报道了显著的成功。Mamba也已整合到图形学习中 (Behrouz and Hashemi, 2024)和模仿学习中 (Correia and Alexandre, 2024)。进一步的进展改善了语言模型,例如加入注意力 (Lieber et al. 2024)、专家混合 (Pióro et al. 2024) 或双向序列处理 (Liang et al., 2024)。

这些进展强调了Mamba的广泛适用性,使其成为多种任务和数据类型的强大且灵活的架构。类似于变换器架构,问题随之而来:Mamba架构是否也可以用于表格问题?虽然Ahamed和Cheng (2024a)已证明Mamba架构可以用于表格数据,但仍需对模型架构进行深入分析以及对表格数据集的优化洞察。论文的贡献可以总结如下:

I. 我们提出了 Mambular,这是对 Mamba 的表格适配,并展示了序列模型在表格问题中的适用性。

II. Mambular 在多个其他竞争的神经网络以及基于树的方法上进行了广泛的基准测试,展示了默认的 Mambular 模型在广泛的数据集上与基于树的模型表现相当或更好。

III. 我们分析了双向处理以及特征交互层对 Mambular 性能的影响,并比较了经典的池化方法。

IV. 最后,我们对 Mambular 的序列性质进行了全面分析,探讨了序列表格模型中特征顺序的影响。

2 方法论

对于表格问题,设

D

=

{

(

x

(

i

)

,

y

(

i

)

)

}

i

=

1

n

\mathcal{D} = {\left\{ \left( {\mathbf{x}}^{\left( i\right) },{y}^{\left( i\right) }\right) \right\} }_{i = 1}^{n}

D={(x(i),y(i))}i=1n​ 为大小为

n

n

n 的训练数据集,

y

y

y 表示可以任意分布的目标变量。每个输入

x

=

(

x

1

,

x

2

,

,

x

J

)

\mathbf{x} = \left( {{x}_{1},{x}_{2},\ldots ,{x}_{J}}\right)

x=(x1​,x2​,…,xJ​) 包含

J

J

J 个特征(变量)。分类和数值特征被区分为

x

(

x

cat

,

x

num

)

\mathbf{x} \equiv \left( {{\mathbf{x}}_{\text{cat }},{\mathbf{x}}_{\text{num }}}\right)

x≡(xcat ​,xnum ​) ,完整的特征向量表示为

x

\mathbf{x}

x 。进一步地,设

x

j

(

c

a

t

)

(

i

)

{x}_{j\left( {cat}\right) }^{\left( i\right) }

xj(cat)(i)​ 表示第

i

i

i 个观测值的第

j

j

j 个分类特征,因此

x

j

(

n

u

m

)

(

i

)

{x}_{j\left( {num}\right) }^{\left( i\right) }

xj(num)(i)​ 表示第

i

i

i 个观测值的第

j

j

j 个数值特征。

Mambular 架构的核心是 FT-Transformer(Gorishniy 等,2021)与 Mamba(Gu 和 Dao,2023)的结合。遵循经典的表格变换器架构,首先对分类特征进行编码和嵌入。与经典语言模型相比,每个分类特征有自己独特的词汇表,以避免二进制或整数编码变量的问题。包括

<

<

< UNK

>

>

> 的标记还可以在训练或推断过程中轻松处理未知或缺失的分类值。

数值特征通过一个简单的线性层映射到嵌入空间。然而,由于单个线性层在信息上并没有超过线性变换,因此采用了 Gorishniy 等(2022)引入的周期线性编码来处理所有数值特征。因此,每个数值特征在通过线性层重新缩放之前首先被编码。简单的决策树用于检测分箱边界,具体使用分类或回归来作为目标依赖编码函数

h

j

(

x

j

(

n

u

m

)

,

y

)

{h}_{j}\left( {{\mathbf{x}}_{j\left( {num}\right) }, y}\right)

hj​(xj(num)​,y) 。设

b

t

{b}_{t}

bt​ 表示来自决策树的决策边界,编码函数在 Eq. 1 PLE 中给出

z

j

(

num

)

t

=

{

0

如果

x

<

b

t

1

,

1

如果

x

b

t

,

x

b

t

1

b

t

2

b

t

1

其他情况.

(1)

{z}_{j\left( \text{ num }\right) }^{t} = \left\{ \begin{array}{ll} 0 & \text{ 如果 }x < {b}_{t - 1}, \\ 1 & \text{ 如果 }x \geq {b}_{t}, \\ \frac{x - {b}_{t - 1}}{{b}_{t - 2} - {b}_{t - 1}} & \text{ 其他情况. } \end{array}\right. \tag{1}

zj( num )t​=⎩

⎧​01bt−2​−bt−1​x−bt−1​​​ 如果 x

Z

\mathbf{Z}

Z 而不是

X

\mathbf{X}

X ,以澄清嵌入与原始特征之间的区别。

图 1:生成输入矩阵并通过 Mamba 块处理。分类特征被标记化并嵌入,类似于经典的语言模型嵌入。数值特征通过一个简单的线性层进行编码和嵌入。Mamba 块的最终输入矩阵是拼接的嵌入

z

R

N

×

J

×

d

\mathbf{z} \in {\mathbb{R}}^{N \times J \times d}

z∈RN×J×d ,其中嵌入维度为

d

d

d 。

随后,这些嵌入通过一系列 Mamba 层共同处理。这包括一维卷积层以及状态空间(SSM)模型(Gu 等,2021;Hamilton,1994)。在通过 SSM 模型之前,特征矩阵的形状为 (BATCH SIZE)

×

J

×

\times \mathrm{J} \times

×J× (EMBEDDING DIMENSION),后面引用为

N

×

J

×

d

N \times J \times d

N×J×d 。重要的是,表格环境中的序列长度指的是变量的数量,因此第二个维度

J

J

J 对应于特征的数量,而不是例如文档的长度。给定矩阵:

A

R

1

×

1

×

d

×

δ

,

B

R

N

×

J

×

1

×

δ

,

Δ

R

N

×

J

×

d

×

1

,

z

R

N

×

J

×

d

×

1

,

\mathbf{A} \in {\mathbb{R}}^{1 \times 1 \times d \times \delta },\;\mathbf{B} \in {\mathbb{R}}^{N \times J \times 1 \times \delta },\;\Delta \in {\mathbb{R}}^{N \times J \times d \times 1},\overline{\mathbf{z}} \in {\mathbb{R}}^{N \times J \times d \times 1},

A∈R1×1×d×δ,B∈RN×J×1×δ,Δ∈RN×J×d×1,z∈RN×J×d×1, 其中

δ

\delta

δ 表示一个内部维度,类似于 Transformer 架构中的前馈维度,

z

\overline{\mathbf{z}}

z 具有与

z

\mathbf{z}

z 相同的元素,但多了一个额外的轴,更新隐状态

h

j

R

N

×

d

×

δ

{\mathbf{h}}_{j} \in {\mathbb{R}}^{N \times d \times \delta }

hj​∈RN×d×δ 的公式为:

h

j

=

exp

(

Δ

3

A

)

:

,

j

,

:

,

:

1

,

2

,

3

h

j

1

+

(

(

Δ

1

,

2

B

)

1

,

2

,

3

z

)

:

,

j

,

:

,

:

.

(2)

{\mathbf{h}}_{j} = \exp {\left( \Delta { \odot }_{3}\mathbf{A}\right) }_{ :, j, : , : }{ \odot }_{1,2,3}{\mathbf{h}}_{j - 1} + {\left( \left( \Delta { \odot }_{1,2}\mathbf{B}\right) { \odot }_{1,2,3}\overline{\mathbf{z}}\right) }_{ :, j, : , : }. \tag{2}

hj​=exp(Δ⊙3​A):,j,:,:​⊙1,2,3​hj−1​+((Δ⊙1,2​B)⊙1,2,3​z):,j,:,:​.(2) 符号

d

{ \odot }_{d}

⊙d​ 表示外积,其中乘法在

d

d

d -th 轴上进行,并在任何单例轴长度与长度为一的轴相遇时并行化

2

{}^{2}

2 。指数函数是逐元素应用的。状态转换矩阵 A 统治着隐状态从上一个时间步到当前时间步的转换,捕捉隐状态如何独立于输入特征演变。输入特征矩阵

B

\mathbf{B}

B 将输入特征映射到隐状态空间,

决定每个特征在每一步如何影响隐状态。门控矩阵

Δ

\mathbf{\Delta }

Δ 作为一个门控机制,调节状态转换和输入特征矩阵的贡献,使模型能够控制上一个状态和当前输入对当前隐状态的影响程度。

图 2: SSM 更新步骤与

h

h

h 的递归更新:隐状态通过遍历序列(特征)进行迭代更新,类似于递归神经网络。最终表示如方程 3-4 中所述生成。

与 FT-Transformer (Gorishniy et al. 2021) 和 TabTransformer (Huang et al., 2020) 相比,Mambular 确实像处理序列一样迭代处理所有变量。因此,特征交互是顺序检测的;并且在第 4 节中分析特征在序列中的位置对性能的影响。此外,需要注意的是,与 TabPFN (Hollmann et al. 2022) 不同,Mambular 不会转置维度,而是对观察进行迭代。因此,可以在大数据集上进行训练,并且可以很好地扩展到任何训练数据大小,正如 Mamba (Gu 和 Dao, 2023) 所做的那样。

在堆叠和进一步处理后,最终表示

x

~

R

N

×

J

×

d

\widetilde{\mathbf{x}} \in {\mathbb{R}}^{N \times J \times d}

x

∈RN×J×d 被检索。在真正的序列数据中,这些是输入标记的上下文化嵌入,对于表格问题,

x

^

\widehat{\mathbf{x}}

x

表示在嵌入空间中的上下文化或特征交互计变量表示。隐状态沿序列维度堆叠形成:

H

=

[

h

0

,

h

1

,

,

h

T

1

]

R

N

×

J

×

d

×

δ

.

\mathbf{H} = \left\lbrack {{\mathbf{h}}_{0},{\mathbf{h}}_{1},\ldots ,{\mathbf{h}}_{T - 1}}\right\rbrack \in {\mathbb{R}}^{N \times J \times d \times \delta }.

H=[h0​,h1​,…,hT−1​]∈RN×J×d×δ. 最终输出表示

x

~

\widetilde{\mathrm{x}}

x

然后通过将堆叠的隐状态与矩阵

C

R

N

×

J

×

1

×

δ

\mathbf{C} \in {\mathbb{R}}^{N \times J \times 1 \times \delta }

C∈RN×J×1×δ 进行矩阵乘法来计算,其中乘法和求和在最后一个轴上进行,并添加向量

α

R

1

×

1

×

d

\alpha \in {\mathbb{R}}^{1 \times 1 \times d}

α∈R1×1×d ,该向量由输入

z

\mathbf{z}

z 进行缩放:

x

~

=

(

H

4

C

)

+

(

α

3

z

)

.

(3)

\widetilde{\mathbf{x}} = \left( {\mathbf{H} \cdot {}_{4}\mathbf{C}}\right) + \left( {\alpha { \odot }_{3}\mathbf{z}}\right) . \tag{3}

x

=(H⋅4​C)+(α⊙3​z).(3) 更明确地,这可以写为:

x

~

i

,

j

,

k

=

δ

H

i

,

j

,

k

,

δ

C

i

,

j

,

1

,

δ

+

α

1

,

1

,

k

z

i

,

j

,

k

.

{\widetilde{x}}_{i, j, k} = \mathop{\sum }\limits_{\delta }{\mathbf{H}}_{i, j, k,\delta }{\mathbf{C}}_{i, j,1,\delta } + {\alpha }_{1,1, k}{\mathbf{z}}_{i, j, k}.

x

i,j,k​=δ∑​Hi,j,k,δ​Ci,j,1,δ​+α1,1,k​zi,j,k​. 其中

C

\mathbf{C}

C 和

α

\alpha

α 是可学习参数。为了最终处理,

x

~

\widetilde{\mathbf{x}}

x

z

{\mathbf{z}}^{\prime }

z′ 逐元素相乘,结果通过一个最终的线性层:

x

~

final

=

(

x

~

1

,

2

,

3

z

)

W

final

+

b

final

.

(4)

{\widetilde{\mathbf{x}}}_{\text{final }} = \left( {\widetilde{\mathbf{x}}{ \odot }_{1,2,3}{\mathbf{z}}^{\prime }}\right) {\mathbf{W}}_{\text{final }} + {\mathbf{b}}_{\text{final }}. \tag{4}

x

final ​=(x

⊙1,2,3​z′)Wfinal ​+bfinal ​.(4) 在传递到最终任务特定模型头之前,沿序列轴可以进行几种池化技术的测试,例如求和池化、平均池化或简单堆叠。

我们进行了一些实验,使用 [CLS] 标记,并分析当仅将 [CLS] 标记传递到任务特定头时模型的性能。考虑到在遍历变量时隐藏状态的递归更新,变量位置也进行了分析,因为变量的顺序可能会影响这个顺序设置。模型通过最小化特定任务损失进行端到端训练,例如,对回归使用均方误差,对分类任务使用分类交叉熵。模型的前向传播概述如图 2 所示。

图 3:模型中单个序列的前向传播。将输入嵌入后,嵌入传递到几个 Mamba 块。表格头是一个单一任务特定输出层。在传递到线性层之前,情境嵌入通过平均池化进行池化。对于双向处理,使用一个反向序列的第二个块,且不可学习矩阵在两个方向之间不共享。

3 实验

Mambular 在多种数据集上进行了基准测试,与当前表现最好的模型(McElfresh et al. 2024)进行对比。FT-Transformer (Gorishniy et al. 2021)、TabTransformer (Huang et al., 2020)、始终表现良好的 XGBoost (Grinsztajn et al., 2022;McElfresh et al., 2024)、基础的多层感知机和 ResNet 是参与此次基准测试的模型。请注意,TabPFN (Hollmann et al. 2022) 未被包含,因为它不适用于较大的数据集。

我们对所有数据集执行 5 折交叉验证,并报告平均结果以及标准差。对于所有神经模型,使用 PLE 编码(公式 1),最大分箱数等于模型维度(大多数模型为 128,包括 MLP 和 ResNet)。所有分类特征均以整数编码。对于回归任务,目标进行了归一化。报导的均方误差(MSE)和曲线下方的面积(AUC)统计数据分别用于回归和分类任务。TabTransformer、FT-Transformer 和 Mambular 使用相同的嵌入架构以及任务特定头,该头由一个没有激活或 dropout 的单一输出层组成。对于 FT-Transformer,使用 [CLS] 标记嵌入进行最终预测,因为研究表明它能提高性能(Thielmann et al. 2024b;Gorishniy et al. 2021)。

所有神经模型使用几个共享参数:开始学习率为

1

e

04

1\mathrm{e} - {04}

1e−04 、权重衰减为 1e-06、相对于验证损失的早期停止耐心为 15 轮、最大训练轮数为 200 轮,以及相对于验证损失的小学习率衰减,衰减因子为 0.1,耐心为 10 轮。此外,使用通用批量大小 128,并针对验证损失返回最佳模型进行测试。对于 TabTransformer、FT-Transformer 和 Mambular,使用相同的嵌入函数。在基准测试中,使用一个简单的 Mambular 架构,利用平均池化,没有特征交互层,也没有双向处理。此外,列/序列始终按数值进行排序。特征首先,其次是分类特征。在这两个组内,特征按原始数据集中提供的顺序排序,数据集的详细信息和预处理可以在附录 A 中找到。模型架构和超参数的详细信息可以在附录 C 中找到。

MambaTab 除了这些流行的表格模型外,我们还测试了 Ahamed 和 Cheng(2024a)提出的架构。MambaTab 是第一个利用 Mamba 块解决表格问题的架构。然而,作者提出使用组合线性层将所有输入投影到单一特征表示中,将特征转变为固定长度为 1 的伪序列。这种方法将方程 2 中的递归更新简化为矩阵乘法,并使得模型由于最后处理中的残差连接而类似于 ResNet。利用序列长度为 1 的序列模型并不能充分利用顺序处理的优势,因为这减少了模型捕获多个特征之间依赖关系的能力。

我们测试了 Ahamed 和 Cheng(2024a)提出的架构,并能够在共享数据集上取得类似的结果,但总体上发现 MambaTab 的表现与 ResNet 相似,符合预期(见表 2 和 23)。此外,我们还尝试转置轴以创建形状为(1)

×

\times

× (BATCH SIZE)

×

\times

× (EMBEDDING DIMENSION)的输入矩阵,如其实现中所述。虽然这种方法借鉴了 TabPFN(Hollmann 等 2022)的想法,但在我们的实验中并未导致性能改进。当使用 PLE 编码并增加层数和维度时,相较于 Ahamed 和 Cheng(2024a)的默认实现,我们能够提高性能。有关 MambaTab 的进一步讨论请见附录 C。

与 XGBoost 的比较 将 mambular 与 XGBoost 进行比较,我们发现,在默认超参数设置下,Mambular 的表现与 XGBoost 一样好,甚至略好于 XGBoost。在 12 个数据集上,Mambular 在 4 个数据集上显著超越 XGBoost,达到

10

%

{10}\%

10% 的显著性水平,而 XGBoost 则在 2 个数据集上超越了 Mambular,达到

10

%

{10}\%

10% 的显著性水平。每个数据集的简单 t 检验的

p

p

p 值已报告,测试以 Gorishniy 等(2021)为基础。当通过 Benjamini-Hochberg(Ferreira 和 Zwinderman, 2006; Benjamini 和 Hochberg, 1995)调整多重检验时,Abalone 结果 - 仅在

10

%

{10}\%

10% 水平上通过标准测试显著 - 不再显著。所有其他结果保持不变

3

{}^{3}

3 。

表格 1:Mambular 与 XGBoost 的比较。左侧表格显示了 5 次折叠的平均 MSE 回归结果。右侧显示(双元)分类结果的平均 AUC 值。在

5

%

5\%

5% 显著性水平上显著更好的值用绿色标注并加粗。在

10

%

{10}\%

10% 显著性水平上显著更好的值下划线标记。数据集的详细信息可以在附录中找到

A

A \uparrow

A↑ 表示越高越好,反之亦然。

结果 总体而言,我们可以确认 Gorishniy 等(2021)提出的 FT-Transformer 架构的强大结果,并由 McElfresh 等(2024)验证。不出所料,XGBoost 在所有数据集上的表现都相当不错,但略微逊色于 FT-Transformer。总体而言,Mambular 在评估的模型中在所有数据集上的平均表现最佳。然而,Mambular 在与 XGBoost 的比较中面临与其他神经方法相似的困难,例如在葡萄酒数据集上。

所有数据集和任务的详细结果如下所示(见表 23 和 24)。原始 MambaTab 实现的结果以及我们结果的讨论可以在附录 C 中找到。有趣的是,对于测试数据集,所有神经模型在葡萄酒数据集上的表现都明显较差(见附录的数据集描述)。类似地,XGBoost 在 Abalone 和 FICO 数据集上的表现不及所有神经模型。此外,我们的发现表明,FT-Transformer 和 Mambular 在具有极少分类特征的数据集上表现很好。

表格 2:模型在所有数据集上的平均排名和标准差。最佳排名用粗体标注。MambaTab 指的是 Ahamed 和 Cheng(2024a)描述的模型。MambaTab

T

{}^{T}

T 指的是转置轴,因此对批量大小进行迭代。MambaTab* 指的是我们的实现。有关 MambaTab 的详细信息,请参见上面的段落。

(例如,FICO、加州住房、Abalone、CPU),尽管利用了原本为离散数据输入设计的嵌入和结构。

表格 3:回归任务的基准测试结果。报告了 5 次折叠的平均均方误差值及其相应的标准差。值越小越好。表现最佳的模型用粗体标注。

表 4:分类任务的基准结果。报告了5折的平均AUC值和相应的标准差。较大的值更好。

3.1 分布回归

为了进一步验证 Mambular 在表格问题上的适用性,我们进行了一个关于分布回归的小任务(Kneib et al. 2023)。分布回归描述了超出均值的回归,即对所有分布参数的建模。因此,位置缩放和形状(LSS)模型可以量化协变量对均值以及对响应假定的潜在复杂分布的任何参数的影响。这些模型的主要优势在于它们能够识别响应分布的所有方面的变化,如方差、偏度和尾部概率,使模型能够正确区分偶然性的不确定性与知识性的不确定性。虽然这在经典统计方法中已经成为一种常见标准(Stasinopoulos 和 Rigby, 2008),但在机器学习社区尚未得到广泛采用。然而,最近的可解释方法(Thielmann et al. 2024a)已经证明了分布回归在表格深度学习中的适用性。此外,像 XGBoostLSS(März, 2019; März 和 Kneib, 2022)这样的模型证明了基于树的模型能够有效地解决此类任务。以下,我们展示 Mambular 的位置缩放和形状(MambularLSS)在最小化负对数似然的同时保持较小均方误差(MSE)时,优于 XGBoostLSS 在连续排名概率评分(CRPS)方面的表现(Gneiting 和 Raftery 2007)。有关 CRPS 指标的简短介绍,请参见附录 D.1。

表 5:阿巴隆和加州住房数据集中对正态分布的分布回归结果。在

5

%

5\%

5% 水平下显著更好的模型用绿色标记。阿巴隆和加州住房在 CRPS 指标上的

p

p

p 值分别为 0.20 和 0.00002。

4 消融

我们测试了 Mambular 的不同架构,以分析不同的池化技术或双向处理是否具有优势或可能降低性能。我们测试了求和池化、最大池化、标准平均池化和仅将序列中的最后一个标记传递给特定任务模型头的最后标记池化。此外,我们评估了可学习的交互层在提高性能方面的有效性。该层通过学习一个交互矩阵

W

R

J

×

J

\mathbf{W} \in {\mathbb{R}}^{J \times J}

W∈RJ×J 来捕捉和建模特征之间的交互,使得交互

=

z

W

= \mathbf{z}\mathbf{W}

=zW ,其中

z

\mathbf{z}

z 是输入特征矩阵,然后通过 SSM 传递。在自然语言处理的 Transformer 网络中,通过 [CLS] 标记嵌入进行池化相当常见(Gorishniy et al. 2021, Thielmann et al., 2024b)。这在表格问题中也被证明是有益的(Thielmann et al. 2024b)。我们不是将 [CLS] 标记附加到每个序列的开头,而是附加到每个序列的末尾。然后执行最后标记池化。

表 6:各种数据集和模型配置的平均 AUC 和平均 MSE。我们测试不同的池化方法、双向处理和可学习的交互层。与默认(平均池化、无交互和无双向处理)相比,显著较差的结果在

5

%

5\%

5% 显著性水平下用红色粗体标记,在

10

%

{10}\%

10% 显著性水平下用下划线和红色标记。所有结果均通过 5 折交叉验证获得,使用与主要结果相同的种子。

有趣的是,我们发现平均池化、无交互和仅单向处理的基础架构在模型配置中始终表现最佳。此外,最后标记池化和 [CLS] 标记池化在四个测试数据集中有两个表现显著较差。进行 5 折交叉验证,所有模型使用相同的超参数。在双向处理过程中,矩阵

A

,

B

\mathbf{A},\mathbf{B}

A,B 和

Δ

\Delta

Δ 不共享,但每个方向都有自己的一组可学习参数。因此,双向模型有额外的可训练参数。

此外,我们分析了序列的顺序和变量在序列中的位置是否影响模型的性能。对于文本数据,单词/标记的顺序混洗具有显著影响,甚至替换单个单词也可能导致完全不同的上下文化嵌入。由于这些上下文化表示在 Mambular 中被池化并直接输入特定任务头,这也可能影响性能。我们测试了两种不同的混洗设置:i)在嵌入层之后混洗嵌入,ii)在通过嵌入层之前混洗变量的序列。所有序列默认按顺序排列,首先是数值特征,随后是分类特征,这与UCI机器学习库中的数据集排列一致。为了进行消融研究,模拟了一个包含5000个样本和10个特征(五个数值特征和五个分类特征)的数据集。数值特征是通过强相关性生成的,包括两个相关性分别为0.8和0.6的特征对。分类特征则创建了四个不同的类别。交互项包括如下内容:两个数值特征之间的交互、一个分类特征与一个数值特征之间的交互,以及两个分类特征之间的交互。在生成目标变量之前,数值特征经过标准化处理。目标变量的构造包括来自每个特征和指定交互项的线性效应,并增加了高斯噪声以增加变异性。我们首先拟合了一个XGBoost模型以进行合理性检查。随后,我们用默认排序(数值特征在前,分类特征在后)、翻转排序以及交换分类和数值特征的排序来拟合Mambular。接着,我们随机打乱顺序并拟合了10个模型。我们发现,即使在这些大型交互和相关效应下,排序对这个模拟数据并没有影响

4

{}^{4}

4 表7:不同特征排序的性能。数值特征以整数表示,分类特征以大写字母表示。数值特征之间的特征交互用蓝色表示。分类特征之间的特征交互用绿色表示,而数值特征与分类特征之间的交互用紫色表示。我们发现,无论是在嵌入层之前还是之后重新排序特征,都不会影响模型的性能。没有一种排序显著优于或劣于默认模型,而所有模型的性能都明显优于XGBoost模型。

为了验证这些结果,我们在四个基准数据集上进行了测试。我们再次发现,在嵌入层之前或之后翻转序列并没有产生显著影响。此外,我们通常发现,序列排序对大多数数据集的模型性能没有显著影响,显著性水平为

5

%

5\%

5% 。

然而,对于加利福尼亚住房数据集,我们发现有一个随机打乱的序列在

5

%

5\%

5% 显著性水平上表现得比基线差。然而,分析该变量序列时,我们发现了一个一致但令人困惑的模式。变量经度和纬度的位置似乎直接影响模型性能。然而,图4显示特征之间的相关性远远强于经度和纬度之间的相关性。此外,我们用简单的线性回归和XGBoost分析了所有成对交互的效应强度。我们发现,经度和纬度之间的交互效应远不如经度和中位收入之间的交互效应重要。我们在附录D中对加利福尼亚住房数据集进行了详细分析。

这一发现虽然仅与四个测试的真实数据集中的一个相关,但表明变量排序可能对Mambular的模型性能产生影响。尽管默认排序

表8:不同特征排序的均值AUC和均值MSE。翻转序列在

5

%

5\%

5% 或

10

%

{10}\%

10% 显著性水平上不会显著影响性能。与默认配置(Num|Cat)在

5

%

5\%

5% 水平上显著不同的值用粗体和红色标记。

图 4:加利福尼亚住房数据集的相关性图

显示出强劲的性能,探索替代排序可能进一步增强 Mambular 的基准结果。正如预期的那样,我们没有发现这些结果适用于基于注意力的模型。有关 FT-Transformer 的结果,请参见附录。

5 限制

所展示的模型在多种数据集上进行了测试,并与一系列模型进行了基准比较。请注意,我们没有进行超参数调优,因为 Grinsztajn et al. (2022) 和 Gorishniy et al. (2021) 的结果表明,大多数模型在没有调优的情况下已经能够良好运行。这些研究表明,虽然超参数调优确实提高了所有模型的性能,但模型的相对排名基本保持不变。这意味着一个在默认配置下表现最佳或最差的模型,即使经过广泛调优,也可能会保持其排名。此外,McElfresh et al. (2024) 发现了类似的结果,进一步加强了超参数调优在大多数模型上具有相似程度的益处而不改变它们的相对性能的观点。

表 9:加利福尼亚住房结果分析

缺乏调优为所有模型留出了改进的空间。然而,对比模型的默认配置在各种研究中经过充分测试,预计如果某个模型可能更受益于超参数调优,那将是 Mambular,因为我们缺乏广泛的文献来指导其默认设置。对于对比模型,我们基于文献选择了参数,以确保其默认设置有意义且表现良好。我们能够重现 Gorishniy et al. (2021) 和 Grinsztajn et al. (2022) 等研究的平均结果。此外,学习率、耐心和epoch数量等关键超参数在所有模型中共享,以实现更一致的方法。所有超参数配置列于附录C中,也可以在 https://github.com/basf/mamba-tabular 找到。

6 结论

我们提出了 Mambular,一种用于表格深度学习的新架构。我们展示了一个真正的序列模型可以应用于表格问题,通过将其视为序列问题来提供对表格数据的新理解和处理方式。我们的结果表明,序列模型在各种数据集中的回归和分类任务中均有效。Mambular 的表现以及其扩展到 MambularLSS 的能力表明它在许多表格任务中的广泛适用性。

尽管 Mamba 与 Transformer 等架构相比仍相对较新,但其迅速被采用表明其在进一步改进方面具有显著潜力。Lieber et al. (2024) 和 Wang et al. (2024) 提出的进展可能对表格应用特别有益。此外,探索特征的最佳排序或通过文本嵌入结合特定列的信息可能进一步增强性能。

将表格数据解释为序列在特征增量学习方面具有显著优势。新特征可以直接附加到序列中,而无需重新训练整个模型。

Copyright © 2088 中国举办世界杯_世界杯足球场地尺寸 - lchjdj.com All Rights Reserved.
友情链接