Революция in-context learning для табличных данных
NVIDIA представила Kumo Tabular, новую архитектуру для работы со структурированными данными. В отличие от традиционных подходов, требующих дообучения (fine-tuning) и настройки гиперпараметров, Kumo Tabular использует механизм in-context learning. Модель принимает контекст из размеченных строк и предсказывает новые значения за один прямой проход (forward pass). Это устраняет необходимость в инженерии признаков и обучении с нуля.
Решение работает через библиотеку structured-data-models (SDM) — GPU-native фреймворк, который также поддерживает TabICLv2, TabFM и KumoRelational. Все модели используют единый интерфейс TableTensor для предобработки и ансамблирования.
Архитектура и обучение на синтетике
Модель построена на базе Transformer, оптимизированного под структуру таблиц. Процесс включает три этапа:
- Cell embedding: Числовые и категориальные значения обрабатываются через обученные Фурье-признаки. Пропуски (missing values) не требуют импутации.
- Row embedding: Используется induced self-attention для линейного роста стоимости по строкам и rotary positions для изучения взаимодействий признаков.
- In-context learning: Финальный Transformer обрабатывает эмбеддинги строк. Ключи и значения контекста вычисляются один раз, что ускоряет инференс.
Важно отметить: Kumo Tabular обучена исключительно на синтетических таблицах, сгенерированных через Structural Causal Models (SCM). Генератор имитирует реальные проблемы: пропуски, высококардинальные категории, тяжелые хвосты распределений и конфликтующие дубликаты. Размеры моделей варьируются от 28M до 215M параметров.
Бенчмарки и сравнение с конкурентами
По данным NVIDIA, Kumo Tabular занимает первое место в рейтинге TabArena с Elo 1950. Модель работает в 17 раз быстрее LimiX-2 на одной RTX 6000 Pro. Ниже приведено сравнение ключевых характеристик с основными конкурентами:
| Характеристика | Kumo Tabular (NVIDIA) | TabICLv2 (Inria) | TabPFN-3 (Prior Labs) | LimiX-2 (Stable AI) | TabFM (Google) |
|---|---|---|---|---|---|
| Параметры | 28M – 215M | ~28M | Не указано | 400M | ~1.64B |
| Лицензия весов | OpenMDW-1.1 (Коммерческая) | BSD-3-Clause | TABPFN-3 License v1.0 (Платная) | Non-Commercial | Non-Commercial |
| Коммерческое использование | Да | Да | Требует лицензии | Нет | Нет |
| Рейтинг TabArena (Elo) | 1950 (1-е место) | Информация недостаточна | Информация недостаточна | Информация недостаточна | Информация недостаточна |
| Скорость (относительная) | 17x быстрее LimiX-2 | Информация недостаточна | Информация недостаточна | База сравнения | Информация недостаточна |
Практическое применение и запуск
Для использования требуется Python 3.11+ и PyTorch 2.7+. Установка осуществляется через pip:
pip install structured-data-modelsПример инициализации модели для классификации:
import sdm
from sklearn.datasets import load_breast_cancer
df = load_breast_cancer(as_frame=True).frame
table = sdm.TableTensor.from_pandas(df).stypes(sdm.infer_stypes(df))
model = sdm.models.KumoTabular(task="classification", device="cuda")
probs = model(x_context=table[:300].drop_columns("target"),
y_context=table[:300]["target"],
x_query=table[300:],
num_estimators=10)Ключевое преимущество Kumo Tabular — сочетание лидерства в бенчмарках с открытой коммерческой лицензией, что делает её привлекательной для внедрения в промышленные решения, где другие open-source альтернативы (TabFM, LimiX-2) запрещены для коммерческого использования.
Источник: MarkTechPost ↗
