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 масиви или pandas DataFrame, включително категорийни характеристики и липсващи стойности. Класификацията поддържа до 10 класа. По подразбиране регресията връща средни предсказания. Версия 1.0.1 добавя медианни и квантилни изходи, базирани на 999 естествени квантила.
Архитектура: усъвършенстване, компресиране, обучение в контекст
Nums AI разделя мрежата на 3 фази: усъвършенстване, компресиране и обучение в контекст. Публикуваните код и конфигурации показват как работи всяка фаза.
Характеристиките се групират в набори по 3. Всяка стойност се вгражда с 16 научени синусоидални и косинусоидални честоти. Липсващите стойности получават собствен научен вектор.
2 етапа на обработка на колоните обобщават всяка група характеристики. Във всеки от тях 128 латентни слота прочитат само обучаващите редове и предават това обобщение на всеки ред. Между 2-та етапа на обработка на колоните етапът на обработка на редовете позволява на групите характеристики да взаимодействат чрез 4 латентни токена. Използва се cross-attention вместо пълно self-attention, което според Nums AI запазва линейната зависимост на разходите от броя характеристики.
След това блок за пулване компресира всеки ред във фиксиран вектор с размерност 512. Към обучаващите редове се добавят етикети. Блок за предсказване с 12 слоя позволява на заявените редове да се насочват към тези етикетирани редове. Заявените редове не могат да променят обучаващия контекст или помежду си.
По подразбиране 8 ансамблови члена споделят едни и същи тегла. Всеки от тях преминава циклично през нормализация none, rank2gaussian, robust или power, с начални стойности за пермутации на характеристиките и класовете.
Резултати в TabArena
Nums AI използва официалния конвейер на TabArena: 51 набора от данни и 816 Full разбиения, с 8 оценителя и начална стойност 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 сгъвания и ограничение от 3000 извадки. Nums AI съобщава, че Causilo се класира на 1-во място по CRPS, R² и RMSE. Поддържащият ScoringBench независимо е проверил резултатите, преди да ги добави.
Скорост и памет
Nums AI също повтори тестовете с 3 модела на 1 H100 GPU с 80 GB памет, като за всяка задача бяха използвани 8 процесорни ядра.
| Модел | Обучение (сек. на 1 хил. реда) | Предсказване (сек. на 1 хил. реда) | Памет на GPU (GiB) |
|---|---|---|---|
| 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])Можете също да изпробвате демо Space в Hugging Face.
Основни изводи
- Causilo има най-високия Elo в TabArena сред единичните модели — общо и по отделните задачи.
- Когато ансамблите от системи бъдат включени, повторното изпълнение от поддържащ поставя модела на 3-то място от 88.
- Смесването на редовете се извършва чрез 4 латентни токена, като така разходите остават линейни спрямо броя характеристики.
- Кодът е под лиценз Apache-2.0; без комерсиален лиценз теглата са предназначени само за изследователски цели.
- Версия 1.0.1 добавя квантилни изходи, така че регресионните интервали работят директно.
Разгледайте хранилището в GitHub и модела в HF. Цялата заслуга е на изследователя на този проект. Също така можете да ни последвате в Twitter и не забравяйте да се присъедините към нашия SubReddit с над 150 хил. членове и да се абонирате за нашия бюлетин. Чакайте! В Telegram ли сте? Вече можете да се присъедините към нас и в Telegram.
Имате нужда да си партнирате с нас за популяризиране на вашето GitHub хранилище ИЛИ страница в Hugging Face ИЛИ представяне на продукт ИЛИ уебинар и т.н.? Свържете се с нас
Преведено автоматично от английски. Оригиналната статия е на връзката по-долу.