Назад к блогу

NVIDIA выпустила Kumo Tabular: базовую модель для табличных данных без обучения под задачу

NVIDIA выпустила Kumo Tabular: базовую модель для табличных данных без обучения под задачу

NVIDIA представила Kumo Tabular — первую базовую модель для табличных данных, которая выдаёт предсказания за один прямой проход, без дообучения, подбора гиперпараметров и ручного конструирования признаков. Подход строится на in-context learning: размеченные строки таблицы подаются как контекст, а модель предсказывает значения в неразмеченных. Это интересно тем, что бросает вызов гегемонии градиентного бустинга в табличных задачах и обещает заменить классический ML-пайплайн одной предобученной моделью.

NVIDIA открыла Kumo Tabular — базовую модель (foundation model) для табличных данных, часть коллекции Kumo Structured. Веса выложены на Hugging Face, код — на GitHub, лицензия — OpenMDW License Agreement версии 1.1, допускающая коммерческое использование.

Модель предобучена только на искусственных таблицах, выпущена в трёх размерах (от 28M до 215M параметров) и выполняет классификацию и регрессию за один прямой проход — без обучения, настройки гиперпараметров и конструирования признаков. По заявлению авторов, она занимает первое место на четырёх бенчмарках: TabArena, BeyondArena, TALENT и ScoringBench.

Как это работает

Таблица подаётся в модель целиком, но разделяется на две части. Строки с известной целевой переменной образуют контекст, строки без неё — запросы, для которых нужно предсказание. В библиотеке structured-data-models (импортируется как sdm) это разделение делается маской по пропускам в целевой колонке:

na_mask = table["target"].isnan()

Данные превращаются в тензор на GPU через sdm.TableTensor.from_pandas, после чего модель sdm.models.KumoTabular вызывается с признаками и целевыми значениями контекста (x_context, y_context) и признаками запроса (x_query).

Внутри финальный Transformer работает на строчных эмбеддингах: контекстные строки взаимодействуют друг с другом, а строки запроса — только с контекстными. Каждое предсказание зависит лишь от контекста и самой строки, а не от того, какие ещё строки оцениваются рядом. Поскольку контекст никогда не смотрит на запросы, его ключи и значения вычисляются один раз и переиспользуются для последующих предсказаний. Голова модели превращает строку запроса в вероятности классов для классификации и в набор квантилей для регрессии.

Классификация и регрессия обучаются как отдельные модели — с кросс-энтропийной потерей для первой и квантильной для второй.

Что даёт in-context learning

Идея заимствована у больших языковых моделей: предобученная модель решает задачу по нескольким примерам, не обновляя ни одного веса. Авторы формулируют это так: модель, предобученная на миллионах таблиц, читает размеченную таблицу как контекст и предсказывает метки новых строк напрямую.

Из того, что контекст никогда не смотрит на запросы, следует практическое свойство: ключи и значения контекста вычисляются один раз и переиспользуются для последующих предсказаний. Строки запроса используют Test-GQA, что уменьшает кэш, читаемый при каждом предсказании.

При росте числа ключей внимание размывается: веса распределяются по всё большему числу элементов, и модель хуже выделяет значимые. Чтобы этого избежать, каждый запрос масштабируется температурой, растущей с логарифмом числа ключей, с отдельным обучаемым коэффициентом для каждой головы внимания. Чем больше ключей, тем сильнее масштабируется запрос, поэтому внимание остаётся резким при увеличении длины или ширины таблицы.

Предыстория: почему таблицы долго оставались за бустингами

Два десятилетия табличные задачи решались градиентно-бустинговыми деревьями, и это работало хорошо. В обзорах, на которые опирается контекст, названы причины, по которым глубокое обучение уступало:

  • Деревья выполняют отбор признаков и эффективно игнорируют неинформативные, разбивая по наиболее релевантным. Глубокие модели к таким признакам не устойчивы, а в табличных данных они встречаются часто.
  • Глубокие модели плохо учатся на нерегулярных функциях, распространённых в табличных данных: граница решения нейросети оказывается заметно более гладкой, чем у дерева.
  • Деревья хорошо работают даже с параметрами по умолчанию, тогда как нейросети требуют больше решений (масштабирование признаков, архитектура) с большим пространством поиска и меньшей переносимостью.
  • Индуктивное смещение деревьев включает ротационно-вариантную процедуру обучения и устойчивость к неинформативным признакам, в отличие от MLP с ротационно-инвариантной процедурой.
  • Градиентный бустинг смещён к объяснению наибольшей доли дисперсии через более простые взаимодействия, с уменьшением вклада по мере роста порядка взаимодействия.

В одном сравнении на 45 наборах данных размером не более 50 000 примеров отмечено, что нейросети могут казаться равными или лучше GBDT, но эффект проявляется только на отдельных подмножествах, где GBDT испытывает трудности. Вывод обзора: при выборе одной модели реализация GBDT остаётся лучшим вариантом, особенно CatBoost.

Почему это сделали

Авторы Kumo Tabular решают проблему трудоёмкого жизненного цикла модели на табличных данных. Каждый новый вопрос означает сбор меток, конструирование признаков, поиск гиперпараметров, валидацию и развёртывание модели, которая ничего не знает о таблицах в целом и учится каждой задаче с нуля.

Альтернатива — базовая модель, которая по размеченной таблице предсказывает метки новых строк за один прямой проход, без обучения, настройки и конструирования признаков, как для классификации, так и для регрессии. Идею авторы заимствуют у больших языковых моделей: предобученная модель решает задачу по нескольким примерам, не обновляя ни одного веса, — это in-context learning, и он применим к таблицам так же, как к тексту.

Что это меняет на практике

Библиотека structured-data-models — GPU-нативная, скачивает веса с Hub при первом использовании и предоставляет препроцессинг, ансамблирование и обработку многих классов. Модель работает только с числовыми и категориальными колонками; текст, изображения и временные метки превращаются в признаки через встроенные рецепты предобработки.

Заявленные метрики:

  • На TabArena — первое место, в 17 раз быстрее LimiX-2 в едином оценочном окружении с одной RTX 6000 Pro.
  • На BeyondArena — первое место.
  • На TALENT — лучший общий ранг по точности классификации, логарифмической потере классификации и RMSE регрессии со средними рангами 6.67, 3.98 и 4.22.
  • На ScoringBench — Kumo Tabular-Large и Medium занимают первое и второе места по среднему рангу.

Ограничения и открытые вопросы

Авторы перечисляют границы применимости. Точность может падать на таблицах, далеко выходящих за диапазоны обучения, или когда строки запроса происходят из другого распределения, чем строки контекста, поэтому перед развёртыванием следует проверять точность и калибровку на собственных отложенных данных.

Один прямой проход покрывает до 10 классов; для любого числа классов библиотека расширяется с помощью error-correcting output codes.

Рецепт обучения и генераторы искусственных данных будут выпущены позже — воспроизводимость пока открытый вопрос. В комментарии сообщества отмечено, что приведённые диагностические результаты получены на одном seed, одном чекпоинте и неравных фиксированных рецептах, и в них нет заявлений о точности, калибровке или общей робастности.

Как обучали

Генератор обучающих данных намеренно воспроизводит несовершенства реальных таблиц: значения пропадают по нескольким шаблонам, часть признаков огрубляется так, что дублирующиеся строки могут расходиться в метке, некоторые категориальные столбцы содержат много уровней, а целевые переменные регрессии могут иметь тяжёлые хвосты. Модель, увидевшая миллионы таких таблиц, учится справляться с ними без очистки данных.

На каждой искусственной таблице модель видит большинство строк с метками как контекст и учится предсказывать метки оставшихся.

Источники

Похожее