HuggingFace镜像/sap-rpt-1-oss
模型介绍
文件和版本
分析

sap-rpt-1-oss

arXiv REUSE status

[!NOTE] 此模型和仓库前身为ConTextTab。尽管代码和仓库现已根据新名称sap-rpt-1-oss进行更新,但模型检查点和功能保持不变。

除了此次开源发布外,您还可以通过 SAP-RPT Playground 免费试用我们最新的商业版SAP‑RPT变体。

描述

论文《ConTextTab: A Semantics-Aware Tabular In-Context Learner》(https://arxiv.org/abs/2506.10707)中所述深度学习模型及推理 pipeline 的实现。

logo

摘要

表格上下文学习(ICL)最近在多个表格预测任务上取得了最先进(SOTA)的性能。 此前,表格上下文学习仅限于小型表格的分类问题,而 TabPFN 和 TabICL 等最新进展已将其应用扩展到更大规模的数据集。 虽然当前原生表格 ICL 架构在结构上高效且能很好地适应表格数据结构,但由于仅在合成数据上进行训练,它们未能充分利用现实世界表格数据中蕴含的丰富语义和世界知识。 另一方面,基于预训练大型语言模型(如 TabuLa-8B)的表格 ICL 模型整合了深度语义理解和世界知识,但由于固有的架构限制,只能利用少量上下文信息。 为了融合这两个领域的优势,我们提出了sap-rpt-1-oss(前称ConTextTab),将语义理解和对齐集成到原生表格 ICL 框架中。 通过为不同数据模态采用专门的嵌入,并在大规模现实世界表格数据上进行训练,我们的模型在一系列广泛的基准测试中与 SOTA 水平具有竞争力,同时在语义丰富的 CARTE 基准测试上树立了新的标准。

引用说明

如果您在研究中使用了本模型,或希望引用我们的研究成果,请按以下格式引用:

@inproceedings{
spinaci2025contexttab,
title={ConTextTab: A Semantics-Aware Tabular In-Context Learner},
author={Marco Spinaci and Marek Polewczyk and Maximilian Schambach and Sam Thelin},
booktitle={Advances in Neural Information Processing Systems (NeurIPS)},
year={2025},
url={https://openreview.net/forum?id=kGMRb4jbTP}
}

要求

本项目使用可从 https://huggingface.co/sap/sap-rpt-1-oss 获取的模型检查点,这些检查点会在运行模型时自动下载。 请注意,下载模型检查点需要登录 Hugging Face。详情请参见说明文档。

Python 3.11 版本的详细要求已在 requirements.txt 文件中列出。

本地开发环境安装:
pip install -e .

从源代码安装:
pip install git+https://github.com/SAP-samples/sap-rpt-1-oss

基本使用

该模型支持分类和回归任务。它接受 pandas DataFrame 或 NumPy 数组形式的输入数据。无需进行预处理,列名和单元格值会通过后台运行的 LLM 自动嵌入,并且能够正确处理任何缺失值。

为获得最佳性能,请使用内存至少为 80 GB 的 GPU,并将上下文大小设置为 8192。对于大型表格,建议将 bagging 因子设置为 8。

若需使用更轻量、速度更快且对 GPU 要求较低的模型,可尝试降低上下文大小(例如设为 2048)并将 bagging 因子设置为 1。

分类

from sklearn.datasets import load_breast_cancer
from sklearn.metrics import accuracy_score
from sklearn.model_selection import train_test_split

from sap_rpt_oss import SAP_RPT_OSS_Classifier

# Load sample data
X, y = load_breast_cancer(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.5, random_state=42)

# Initialize a classifier, 8k context and 8-fold bagging gives best performance, reduce if running out of memory
clf = SAP_RPT_OSS_Classifier(max_context_size=8192, bagging=8)

clf.fit(X_train, y_train)

# Predict probabilities
prediction_probabilities = clf.predict_proba(X_test)
# Predict labels
predictions = clf.predict(X_test)
print("Accuracy", accuracy_score(y_test, predictions))

回归

from sklearn.datasets import fetch_openml
from sklearn.metrics import r2_score
from sklearn.model_selection import train_test_split

from sap_rpt_oss import SAP_RPT_OSS_Regressor

# Load sample data
df = fetch_openml(data_id=531, as_frame=True)
X = df.data
y = df.target.astype(float)

# Train-test split
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.5, random_state=42)

# Initialize the regressor, 8k context and 8-fold bagging gives best performance, reduce if running out of memory
regressor = SAP_RPT_OSS_Regressor(max_context_size=8192, bagging=8)

regressor.fit(X_train, y_train)

# Predict on the test set
predictions = regressor.predict(X_test)

r2 = r2_score(y_test, predictions)
print("R² Score:", r2)

已知问题

无已知问题

如何获取支持

如果您发现错误或对内容有疑问,请在本代码库中创建issue。

贡献

如果您希望贡献代码、提供修复或改进,请提交拉取请求。出于法律原因,贡献者在向本项目创建第一个拉取请求时,将被要求接受开发者贡献协议(DCO)。此过程会在提交过程中自动进行。SAP采用Linux基金会的标准DCO文本。

许可

版权所有 (c) 2025 SAP SE 或其关联公司。保留所有权利。本项目根据Apache软件许可协议2.0版授权,除非LICENSE文件中另有说明。

模型检查点是在the T4 dataset上训练的,而该数据集又是the TabLib dataset的一个子集。因此,它们继承了其中描述的相同限制,特别是仅用于研究目的。