Nums AI выпустила Causilo: табличная базовая модель, возглавившая TabArena среди одиночных моделей
Nums AI выпустила Causilo — предварительно обученную базовую табличную модель для классификации и регрессии. Causilo поставляется с интерфейсом scikit-learn, кодом по лицензии Apache-2.0 и предварительно обученными весами на Hugging Face. В TabArena она имеет самый высокий Elo среди одиночных моделей как для классификации, так и для регрессии.
Можно ли её развернуть? Да, уже сегодня — для исследований и оценки, на CUDA или CPU. Для коммерческого использования, эксплуатации в продакшене и доступа через размещённый API требуется отдельная лицензия от Nums AI.
Что делает Causilo
Causilo — это модель обучения в контексте. Вызов fit не обновляет предварительно обученные веса. Она сохраняет обучающие строки в качестве контекста и предсказывает значения для строк запроса за один прямой проход. Согласно своей заявке в TabArena, Nums AI предварительно обучала Causilo только на синтетических данных, без использования наборов данных TabArena.
Входными данными могут быть массивы NumPy или DataFrame pandas, включая категориальные признаки и пропущенные значения. Классификация поддерживает до 10 классов. По умолчанию регрессия возвращает средние предсказания. Версия 1.0.1 добавляет медианные и квантильные выходные данные на основе 999 встроенных квантилей.
Архитектура: уточнение, сжатие и обучение в контексте
Nums AI разделяет сеть на 3 этапа: уточнение, сжатие и обучение в контексте. Выпущенный код и конфигурации показывают, как работает каждый этап.
Признаки объединяются в группы по 3. Каждое значение встраивается с помощью 16 обучаемых синусоидальных и косинусоидальных частот. Для пропущенных значений используется собственный обучаемый вектор.
2 этапа обработки столбцов суммируют информацию по каждой группе признаков. На каждом из них 128 латентных слотов считывают только обучающие строки и передают это резюме каждой строке. Между 2 этапами обработки столбцов этап обработки строк позволяет группам признаков взаимодействовать через 4 латентных токена. Вместо полноценного self-attention используется cross-attention, что, по словам Nums AI, сохраняет линейную зависимость затрат от числа признаков.
Затем блок пулинга сжимает каждую строку в вектор фиксированного размера — 512 измерений. К обучающим строкам добавляются метки. Блок предсказания из 12 слоёв позволяет строкам запроса обращаться к этим размеченным строкам. Строки запроса не могут изменять обучающий контекст или друг друга.
По умолчанию 8 участников ансамбля используют одни и те же веса. Каждый из них циклически применяет одно из преобразований: none, rank2gaussian, robust или power normalization, с фиксированными перестановками признаков и классов.
Результаты в TabArena
Nums AI использовала официальный конвейер TabArena: 51 набор данных и 816 полных разбиений, с 8 оценщиками и seed 42. Сопровождающий TabArena повторно провёл полную оценку и получил тот же общий Elo — 1794.
| Задача | Elo Causilo | Следующая лучшая одиночная модель | Потенциал улучшения Causilo |
|---|---|---|---|
| Общий результат | 1792.9 | TabFM, 1764.4 | 0.0684 |
| Классификация | 1771.8 | EXAONE Tabular, 1758.8 | 0.0875 |
| Регрессия | 2032.6 | TabFM, 1992.8 | 0.0125 |
В число участников входят TabFM от Google Research, EXAONE Tabular от LG AI Research и TabPFN-3 от Prior Labs (1636.2 в общем зачёте).
При чтении этих показателей стоит учитывать несколько моментов:
- Позиции №1 не учитывают системные решения. Если включить системы, повторный запуск сопровождающего проекта поставил Causilo на 3-е место из 88 в общем зачёте.
- По потенциалу улучшения TabFM по-прежнему лидирует в общем зачёте и в классификации. Causilo лидирует в регрессии.
- Доверительные интервалы Elo у лидеров перекрываются, поэтому преимущество над TabFM и EXAONE Tabular невелико.
- Nums AI также указывает Xiaomi-TabLDM и Mitra-v2 от Amazon ниже Causilo. Ни одна из этих моделей не представлена в файлах бенчмарка в репозитории Causilo.
Результаты ScoringBench
ScoringBench оценивает регрессионные модели с помощью корректных правил скоринга, таких как CRPS, а также RMSE и R². Nums AI представила Causilo 1.0.1 на 101 наборе данных, по 5 фолдов на каждый, с ограничением в 3 000 образцов. Nums AI сообщает, что Causilo занимает 1-е место по CRPS, R² и RMSE. Сопровождающий ScoringBench независимо проверил результаты перед их публикацией.
Скорость и память
Nums AI также повторно протестировала 3 модели на 1 GPU H100 с 80 ГБ памяти, выделив по 8 ядер CPU на задачу.
| Модель | Обучение (с на 1 тыс. строк) | Предсказание (с на 1 тыс. строк) | Память GPU (ГиБ) |
|---|---|---|---|
| Causilo | 2.504 | 0.251 | 8.15 |
| TabICLv2 | 3.449 | 0.303 | 8.37 |
| TabPFN-3 | 4.18 | 0.686 | 0.88 |
В этом тесте Causilo быстрее всего выполняет как обучение, так и предсказание. TabPFN-3 использует значительно меньше памяти GPU. Установка use_kv_cache=True переносит обработку контекста на этап обучения, увеличивая расход памяти для ускорения повторных предсказаний.
Начало работы
Для Causilo требуются Python от 3.10 до 3.12 и PyTorch 2.13 или новее. При первом обучении контрольная точка загружается автоматически.
# pip install causilo
from causilo import CausiloClassifier, CausiloRegressor
clf = CausiloClassifier(n_estimators=8, random_state=42)
clf.fit(X_train, y_train)
proba = clf.predict_proba(X_test)
reg = CausiloRegressor()
reg.fit(X_train, y_train)
bands = reg.predict(X_test, output_type="quantiles", quantiles=[0.05, 0.5, 0.95])Также можно попробовать демонстрационное пространство на Hugging Face.
Ключевые выводы
- Causilo имеет самый высокий Elo среди одиночных моделей в TabArena — как в общем зачёте, так и по отдельным задачам.
- С учётом системных ансамблей повторный запуск сопровождающего проекта ставит её на 3-е место из 88.
- Смешивание строк происходит через 4 латентных токена, что сохраняет линейную зависимость затрат от числа признаков.
- Код распространяется по лицензии Apache-2.0; веса предназначены только для исследований без коммерческой лицензии.
- Версия 1.0.1 добавляет квантильные выходные данные, поэтому интервалы регрессии работают сразу после установки.
Посетите репозиторий на GitHub и модель на HF. Вся благодарность исследователю этого проекта. Также подписывайтесь на нас в Twitter и не забудьте присоединиться к нашему сабреддиту по машинному обучению с более чем 150 тыс. участников и подписаться на нашу рассылку. Стоп! Вы пользуетесь Telegram? теперь к нам можно присоединиться и в Telegram.
Хотите сотрудничать с нами для продвижения вашего репозитория GitHub, страницы на Hugging Face, выпуска продукта, вебинара и т. д.? Свяжитесь с нами
Переведено автоматически с английского. Оригинал статьи — по ссылке ниже.