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другая модель для табличных данных, с которой сравнивают Kumo Tabular в едином оценочном окружении с одной RTX 6000 Pro.
- На BeyondArena — первое место.
- На TALENT — лучший общий ранг по точности классификации, логарифмической потере классификации и RMSEсреднеквадратичная ошибка, корень из среднего квадрата отклонений предсказаний от истинных значений регрессии со средними рангами 6.67, 3.98 и 4.22.
- На ScoringBench — Kumo Tabular-Large и Medium занимают первое и второе места по среднему рангу.
Ограничения и открытые вопросы
Авторы перечисляют границы применимости. Точность может падать на таблицах, далеко выходящих за диапазоны обучения, или когда строки запроса происходят из другого распределения, чем строки контекста, поэтому перед развёртыванием следует проверять точность и калибровку на собственных отложенных данных.
Один прямой проход покрывает до 10 классов; для любого числа классов библиотека расширяется с помощью error-correcting output codesсхема кодирования классов наборами битовых меток, при которой модель решает несколько двоичных задач, а исходный класс восстанавливается по совокупности их ответов.
Рецепт обучения и генераторы искусственных данных будут выпущены позже — воспроизводимость пока открытый вопрос. В комментарии сообщества отмечено, что приведённые диагностические результаты получены на одном seed, одном чекпоинте и неравных фиксированных рецептах, и в них нет заявлений о точности, калибровке или общей робастности.
Как обучали
Генератор обучающих данных намеренно воспроизводит несовершенства реальных таблиц: значения пропадают по нескольким шаблонам, часть признаков огрубляется так, что дублирующиеся строки могут расходиться в метке, некоторые категориальные столбцы содержат много уровней, а целевые переменные регрессии могут иметь тяжёлые хвосты. Модель, увидевшая миллионы таких таблиц, учится справляться с ними без очистки данных.
На каждой искусственной таблице модель видит большинство строк с метками как контекст и учится предсказывать метки оставшихся.