pyg-team/pytorch-frame

Tabular Deep Learning Library for PyTorch

解决的问题

PyTorch Frame 为包含文本、图像和数值等混合列类型的异构表格数据提供了一个模块化框架,用于构建深度学习模型。它克服了传统树模型(如 GBDT)的局限性,支持与下游模型更好地集成,并能够处理文本、图像和时间戳等复杂列类型,同时兼顾数值和分类数据。

工作原理

该库采用由三个主要组件构成的模块化架构:

  1. Materialization:将原始的 pandas DataFrame 转换为适合 PyTorch 训练的 TensorFrame 格式。
  2. FeatureEncoder:将各种列类型编码为统一大小的隐藏嵌入。
  3. TableConv:建模这些列嵌入之间的交互。
  4. Decoder:对嵌入进行池化,生成每行的最终预测或嵌入。

支持与 LLM 嵌入 API(如 OpenAI、Cohere 和 Voyage AI)以及 Hugging Face transformers 集成,用于编码文本数据。

适用人群

专为希望实验深度表格模型或使用标准化、模块化方法与 GBDT 进行性能对比的深度学习研究人员和实践者设计。

主要亮点

  • 多样列类型支持:支持数值、分类、多分类、文本(嵌入或分词)、时间戳和图像。
  • 预实现模型:包含 Trompt、FTTransformer、TabNet、ExcelFormer 和 TabTransformer 等前沿模型。
  • PyTorch 生态系统集成:可无缝集成其他 PyTorch 库,包括用于关系数据库学习的 PyG。
  • 基准数据集:提供一系列即用型表格数据集和基准测试工具,用于比较深度学习模型与 GBDT 的性能。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目