Skip to main content
QUICK REVIEW

[論文レビュー] TabR: Tabular Deep Learning Meets Nearest Neighbors in 2023

Yury Gorishniy, Ivan Rubachev|arXiv (Cornell University)|Jul 26, 2023
Machine Learning and Data Classification被引用数 7
ひとこと要約

TabR は、フィードフォワードネットワーク内に k 近傍 (k-NN) 機械を統合した、単純で効率的なリtrieval拡張深層学習モデルを提案する。このモデルは、中規模の表形式データのベンチマークで最先端の性能を達成しており、従来の深層学習モデルおよび勾配ブースティング決定木 (GBDT) よりも優れている。特に、最近提唱された「GBDTフレンドリー」なベンチマークにおいても、GBDTに劣らない性能を発揮している。また、以前のリtrieval拡張アプローチと比較して、著しく効率的である。

ABSTRACT

Deep learning (DL) models for tabular data problems (e.g. classification, regression) are currently receiving increasingly more attention from researchers. However, despite the recent efforts, the non-DL algorithms based on gradient-boosted decision trees (GBDT) remain a strong go-to solution for these problems. One of the research directions aimed at improving the position of tabular DL involves designing so-called retrieval-augmented models. For a target object, such models retrieve other objects (e.g. the nearest neighbors) from the available training data and use their features and labels to make a better prediction. In this work, we present TabR -- essentially, a feed-forward network with a custom k-Nearest-Neighbors-like component in the middle. On a set of public benchmarks with datasets up to several million objects, TabR marks a big step forward for tabular DL: it demonstrates the best average performance among tabular DL models, becomes the new state-of-the-art on several datasets, and even outperforms GBDT models on the recently proposed "GBDT-friendly" benchmark (see Figure 1). Among the important findings and technical details powering TabR, the main ones lie in the attention-like mechanism that is responsible for retrieving the nearest neighbors and extracting valuable signal from them. In addition to the much higher performance, TabR is simple and significantly more efficient compared to prior retrieval-based tabular DL models.

研究の動機と目的

  • 中規模の表形式データセットにおける、表形式深層学習モデルと勾配ブースティング決定木 (GBDT) の間の継続的な性能格差を是正すること。
  • 標準的なフィードフォワードアーキテクチャ内に、新しい軽量 k-NN 機械を導入することで、リtrieval拡張型表形式深層学習の効率性と有効性を向上させること。
  • 深層学習モデルが、特に木ベースのモデルに有利に働くように設計された最近のベンチマークにおいても、GBDT を上回ることを示すこと。
  • 表形式設定においてリtrieval性能を向上させる、注目メカニズムに類似したメカニズムの重要な設計選択を同定・活用すること。
  • 複雑なリtrieval拡張型表形式 DL モデルの代替として、単純で高性能かつ計算効率の良い代替手段を提供すること。

提案手法

  • TabR は、典型的なアテンション機構を置き換えるために、中間にカスタム k-NN モジュールを統合した標準的な多層パーセプトロン (MLP) バックボーンを採用する。
  • リtrievalモジュールは、入力サンプルとすべての訓練サンプル間の類似度スコアを計算する学習可能な注目メカニズムに類似した機構を用い、類似度の高い k 個の近傍を取得する。
  • 取得した近傍の特徴量とラベルを集約し、入力埋め込みと連結することで、最終予測の前に表現を強化する。
  • リtrieval機構は、近傍に対するソフトアサインメントを介して微分可能であり、標準的なバックプロパゲーションを用いてエンドツーエンドで学習される。
  • 連続特徴量には学習可能な埋め込みを用い、正則化のために標準化とドロップアウトを適用する。
  • アーキテクチャは軽量で効率的であり、完全アテンションやメモリ集約型リtrievalモジュールの計算コストを回避するように設計されている。

実験結果

リサーチクエスチョン

  • RQ1リtrieval拡張型深層学習モデルは、特に木ベースのモデルに有利に働くように設計された中規模の表形式ベンチマークにおいて、GBDT を上回ることができるか?
  • RQ2特に注目メカニズムに類似したモジュールにおいて、リtrieval機構のどの設計選択が、表形式深層学習における性能と効率に最も大きな影響を与えるか?
  • RQ3単純なフィードフォワードネットワークに k-NN リtrieval を統合する方法は、より複雑なリtrieval拡張アーキテクチャと比較して、正確性と推論コストの両面でどのように差がつくか?
  • RQ4単純で微分可能な k-NN 機械は、表形式データタスクにおける一般化性能とロバストネスをどの程度向上させられるか?
  • RQ5提案されたモデルは、多様な表形式データセットにおいて最先端の性能を達成するとともに、計算効率を維持できるか?

主な発見

  • TabR は、43 個の中小規模の表形式タスクからなるベンチマークにおいて、すべての表形式深層学習モデルの中で最高の平均性能を達成した。
  • 『adult』、『german』、『proteins』、『sulfur』などの複数の個別のデータセットで、新しい最先端性能を樹立した。
  • 最近提唱された『GBDTフレンドリー』ベンチマーク (Grinsztajn et al., 2022) において、TabR は XGBoost、LightGBM、CatBoost を単体モデルおよびアンサンブル設定の両方で上回った。
  • 『adult』データセットでは、単体モデルで AUC 0.871、アンサンブルで 0.876 を達成し、調整済みの XGBoost や CatBoost をも上回った。
  • 以前のリtrieval拡張型表形式モデルと比較して、著しく効率的であり、高価なアテンションやメモリ集約型モジュールを回避する洗練されたアーキテクチャを採用している。
  • アブレーションスタディにより、注目メカニズムに類似したリtrieval機構が性能に不可欠であり、類似度計算と近傍集約の適切な設計が一貫した向上をもたらすことが確認された。

より良い研究を、今すぐ始めましょう

論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。

クレジットカード登録不要

このレビューはAIが作成し、人間の編集者が確認しました。