
从傅里叶特征到FLASH注意力TabFM核心代码实现原理深度解析【免费下载链接】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 开源的零样本表格预测模型无需训练即可完成表格数据的分类与回归。本文深度解析 TabFM 核心代码实现从 CellEmbedder 的傅里叶特征Fourier Features编码到省内存的 FLASH 注意力Flash Attention机制带你完整读懂这套 Transformer 架构是如何被改造成「表格引擎」的。一、TabFM 是什么免训练的智能表格模型与「每个数据集都要从头训练」的传统机器学习不同TabFM 的核心卖点是in-context learning上下文学习推理时不需要在你的数据集上训练任何参数模型把你的训练行当作「上下文」读入直接对新的测试行做即时预测API 兼容 scikit-learn提供 TabFMClassifier 与TabFMRegressor两个估计器混合数值/类别列开箱即用支持 JAX 与 PyTorch 双后端可在 examples/classification_example.py、examples/regression_example.py 中直接运行。一句话理解你把带标签的训练表喂给它它像读文档一样「读表」然后回答测试行的标签。二、整体架构一张表如何变成预测结果模型主体定义在 TabFM数据流经三个依次嵌套的阶段阶段模块作用1. 列向嵌入CellEmbedder ColEmbedding把每个「单元格」变成向量并让同一列内的分布信息互相感知2. 行向交互RowInteraction用 Transformer 编码器捕捉一行内不同特征之间的交互3. 数据集级 ICLICLearning12 层v1.0.0 为 24 层Transformer 读取「训练行上下文」对测试行输出预测v1.0.0 预训练权重对应的完整超参固化在 Configembed_dim256、swiglu激活、feature_groupTrue3 个相邻特征一组、傅里叶特征开启32 个频率。# 数据流概览形状记号见 model.py 文件头注释 X(B,T,H) ──▶ 单元格嵌入(B,T,H,E) ──▶ 列嵌入 ──▶ CLS ──▶ 行交互 ──▶ ICL ──▶ 预测 └ 傅里叶特征在这里 └ 诱导注意力在这里 └ KV缓存在这里三、傅里叶特征CellEmbedder 如何读懂每一个单元格一个单元格只是一个标量如年龄 35、职业 manager 编码后的 1。直接线性投影很难区分 34 和 35而傅里叶特征把标量映射到高频振荡空间让模型获得更丰富的「频谱」表示。CellEmbedder 中初始化了两组可学习频率银行数值列与类别列分开避免相互干扰# 频率 bank形状 (in_dim, 32)标准正态 × sigma(1.0) 初始化 self.fourier_frequencies nnx.Param(normal(...) * sigma) # 数值特征 self.fourier_frequencies_cat nnx.Param(normal(...) * sigma) # 类别特征前向传播只有三步源码x_proj jnp.einsum(...i,if-...if, X, fourier_frequencies) # 值 × 频率 feats jnp.concatenate([jnp.sin(x_proj), jnp.cos(x_proj)], -1) # 64维频谱 emb self.in_linear(feats) # 投影到 256 维两个值得注意的设计细节类别列走独立的频率 bank注释明确指出「数值特征需要度量保持的中等频率类别特征需要高频来解耦相邻的整数编码」且数值/类别各有独立线性头按cat_mask逐槽选择特征分组开启feature_group后feature_grouping 用(2^i)-1的偏移量让每个单元格「看到」自身及右侧 1、3 个邻居类似滑窗增强局部特征关联。此外回归模型会把标签 y 嵌入后只加到训练行的单元格上ADD_Y_TO_X_POST_EMBEDDING让列分布天然携带目标信息。四、FLASH 注意力不落地 T×T 矩阵的省内存实现标准注意力要显式构造 T×T 的权重矩阵当上下文行很多时内存开销巨大。TabFM 通过 AttentionImplementation 枚举支持四种实现jax/jax_vmap_on_head_dim/flash/none其中FLASH分支直接调用仓库内置的内存高效注意力调用处if self.attention_impl AttentionImplementation.FLASH: attn_output memory_efficient_attention.dot_product_attention_multihead( queryq, keyk, valuev, biasattention_bias, query_chunk_size128, key_chunk_size128) # 按128分块实现位于 memory_efficient_attention.py核心是在线 softmaxonline softmax 分块扫描分块Query 和 Key/Value 各自切成 128 长度小块_memory_efficient_attention永远不在显存中生成完整 T×T 矩阵增量归约用 _AttentionSummary 三元组加权分子、指数和分母、历史最大值沿 Key 块用lax.scan逐步累加数值稳定每处理一个新块在 _summarize_chunk 里更新全局最大值并乘以修正项exp(old_max - new_max)等价于完整 softmax 却只需 O(T) 中间内存。这也解释了 prefill 中把序列长度补齐到 128 倍数的代码——正好对应 FLASH 的block_size。默认加载权重时的注意力配置load 函数很有讲究col_attention_implflash, # 列嵌入序列最长全部行用FLASH省内存 row_attention_impljax, # 行交互序列只是特征数普通注意力即可 icl_attention_implflash # ICL长上下文同样用FLASH五、其他关键机制速览RoPE 旋转位置编码RotaryEmbedding 只对行编码器启用rope_base100000让模型感知「特征列的相对位置」ICL 编码器则不启用因为行的先后顺序没有语义。诱导注意力Induced AttentionInducedSelfAttentionBlock 是列嵌入的「加速器」先用 256 个可学习诱导点压缩全表阶段1诱导点←全部行再让全部行读取诱导点阶段2把列内自注意力的复杂度从 O(T²) 降到 O(T·I)。QK 归一化与逐维缩放MultiheadAttention 对 Q/K 做 RMSNorm并用 PerDimScale 的 softplus 缩放训练更稳定。KV 缓存的 prefill/decode 两段式推理先prefill()读入全部训练行并缓存 KV再对测试行逐批decode()避免重复计算上下文非常适合大表推理。注意力掩码防泄漏ICLearning 中 train_mask 保证任何测试行只能 attend 到训练行测试样本之间互相「不可见」。六、动手跑起来5 分钟验证核心链路git clone https://gitcode.com/gh_mirrors/ta/tabfm cd tabfm pip install -e .[jax]import pandas as pd from tabfm import TabFMClassifier from tabfm import tabfm_v1_0_0_jax as tabfm_v1_0_0 clf TabFMClassifier(modeltabfm_v1_0_0.load()) X pd.DataFrame({age: [25.0, 45.0], job: [a, b]}) clf.fit(X, y[low, high]) print(clf.predict(X)) # 傅里叶特征→FLASH注意力→ICL一条链路走完结语为什么这套代码值得精读TabFM 把三件「经典 Transformer 技巧」精准地嫁接到表格场景傅里叶特征解决标量单元格表达力不足FLASH 内存高效注意力解决长上下文的内存瓶颈诱导注意力 CLS 聚合 KV 缓存则把行、列两个维度都做了计算裁剪。如果你正在做表格数据或 Foundation Model 方向建议按本文顺序精读这三个源文件架构与傅里叶特征tabfm/src/jax/model.pyFLASH 注意力实现tabfm/src/jax/memory_efficient_attention.pyv1.0.0 权重加载tabfm/src/jax/tabfm_v1_0_0.py理解之后你会发现所谓「表格基础模型」本质就是一座带缓存、带频域编码、为二维数据量身裁剪过的 Transformer。【免费下载链接】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),仅供参考