pyg-team/pytorch-frame
Tabular Deep Learning Library for PyTorch
解决的问题
PyTorch Frame 为包含文本、图像和数值等混合列类型的异构表格数据提供了一个模块化框架,用于构建深度学习模型。它克服了传统树模型(如 GBDT)的局限性,支持与下游模型更好地集成,并能够处理文本、图像和时间戳等复杂列类型,同时兼顾数值和分类数据。
工作原理
该库采用由三个主要组件构成的模块化架构:
- Materialization:将原始的 pandas DataFrame 转换为适合 PyTorch 训练的
TensorFrame格式。 - FeatureEncoder:将各种列类型编码为统一大小的隐藏嵌入。
- TableConv:建模这些列嵌入之间的交互。
- Decoder:对嵌入进行池化,生成每行的最终预测或嵌入。
支持与 LLM 嵌入 API(如 OpenAI、Cohere 和 Voyage AI)以及 Hugging Face transformers 集成,用于编码文本数据。
适用人群
专为希望实验深度表格模型或使用标准化、模块化方法与 GBDT 进行性能对比的深度学习研究人员和实践者设计。
主要亮点
- 多样列类型支持:支持数值、分类、多分类、文本(嵌入或分词)、时间戳和图像。
- 预实现模型:包含 Trompt、FTTransformer、TabNet、ExcelFormer 和 TabTransformer 等前沿模型。
- PyTorch 生态系统集成:可无缝集成其他 PyTorch 库,包括用于关系数据库学习的 PyG。
- 基准数据集:提供一系列即用型表格数据集和基准测试工具,用于比较深度学习模型与 GBDT 的性能。
相关
- 项目
- 项目
- 项目
- 项目
- 项目