Sakhanda Wire
NVDA $230.86 +1.09% MSFT $512.80 -0.02% GOOGL $338.24 -1.70% META $725.93 +0.10% AMZN $248.23 -0.37%
← Към новините

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.9TabFM, 1764.40.0684
Класификация1771.8EXAONE Tabular, 1758.80.0875
Регресия2032.6TabFM, 1992.80.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)
Causilo2.5040.2518.15
TabICLv23.4490.3038.37
TabPFN-34.180.6860.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 ИЛИ представяне на продукт ИЛИ уебинар и т.н.? Свържете се с нас

Преведено автоматично от английски. Оригиналната статия е на връзката по-долу.

Първоначално публикувано от MarkTechPost на

Прочетете оригинала в MarkTechPost ↗

Текстът и изображенията са собственост на MarkTechPost и са възпроизведени тук с посочване на авторството и връзка към оригиналната публикация.

← Към новините

Още новини

Всички последни новини