[논문 리뷰] TabR: Tabular Deep Learning Meets Nearest Neighbors in 2023
TabR는 피드포워드 네트워크 내부에 k-최근접 이웃(k-NN) 메커니즘을 통합한 간단하고 효율적인 검색 증강 딥러닝 모델을 제안한다. 이 모델은 표준 데이터베이스에서 최신 기준으로 평가된 'GBDT 우호적' 벤치마크에서도 기존 딥러닝 모델과 기울기 부스팅 결정 트리(GBDT)를 모두 능가하는 최고 성능을 기록했으며, 이전의 검색 증강 접근 방식보다 훨씬 더 효율적이다.
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 메커니즘을 표준 피드포워드 아키텍처에 통합하여, 검색 증강 표준 딥러닝의 효율성과 효과를 향상시키기 위해.
- 특히 트리 기반 모델에 유리하게 설계된 최근 제안된 벤치마크에서 딥러닝 모델이 GBDT를 능가할 수 있음을 입증하기 위해.
- 표준 데이터 설정에서 검색 성능을 향상시키는 데 기여하는 어텐션 유사 메커니즘의 핵심 설계 요소를 식별하고 활용하기 위해.
- 복잡한 검색 증강 표준 딥러닝 모델의 대안으로 단순하고 높은 성능를 발휘하며 계산 비용이 낮은 모델을 제공하기 위해.
제안 방법
- 표준 다층 퍼셉트론(MLP) 기반 아키텍처에 특수한 k-NN 유사 구성 요소를 중간에 삽입하여 일반적인 어텐션 메커니즘을 대체한다.
- 검색 구성 요소는 입력 샘플과 모든 훈련 샘플 간의 유사도 점수를 계산하기 위해 학습 가능한 어텐션 유사 메커니즘을 사용하며, 가장 유사한 k개의 이웃을 검색한다.
- 검색된 이웃의 특징 및 레이블을 집계하여 입력 임bedding에 연결함으로써 예측 이전에 표현을 풍부하게 한다.
- 표준 역전파를 사용하여 엔드 투 엔드로 학습되며, 이웃에 대한 소프트 할당을 통해 검색 메커니즘이 미분 가능하다.
- 연속형 특징에 대해 학습 가능한 임베딩을 사용하고, 정규화 및 드롭아웃을 적용하여 정규화를 수행한다.
- 아키텍처는 경량화되어 있으며, 전체 어텐션 또는 메모리 집약적인 검색 모듈의 계산 오버헤드를 피한다.
실험 결과
연구 질문
- RQ1검색 증강 딥러닝 모델이 특히 트리 기반 모델에 유리하게 설계된 중규모 표준 데이터 벤치마크에서 GBDT를 능가할 수 있는가?
- RQ2특히 어텐션 유사 구성 요소에서의 검색 메커니즘 설계 선택 사항 중 성능과 효율성에 가장 큰 영향을 미치는 요소는 무엇인가?
- RQ3간단한 피드포워드 네트워크에 k-NN 검색을 통합하는 방식이 더 복잡한 검색 증강 아키텍처와 비교해 정확도와 추론 비용 측면에서 어떻게 다른가?
- RQ4간단하고 미분 가능한 k-NN 메커니즘이 표준 데이터 작업에서 일반화 능력과 견고성을 얼마나 향상시킬 수 있는가?
- RQ5제안된 모델이 다양한 표준 데이터셋에서 최고 성능를 달성하면서도 계산 비용이 낮은가?
주요 결과
- TabR는 43개의 중규모 표준 데이터 작업으로 구성된 벤치마크에서 모든 표준 딥러닝 모델 중 평균 성능이 가장 뛰어나다.
- 특히 'adult', 'german', 'proteins', 'sulfur' 데이터셋에서 새로운 최고 성능를 기록했다.
- 최근 제안된 'GBDT 우호적' 벤치마크(Grinsztajn et al., 2022)에서 단일 모델 및 앙상블 설정 모두에서 XGBoost, LightGBM, CatBoost를 능가했다.
- 'adult' 데이터셋에서 TabR는 단일 모델 기준 AUC 0.871, 앙상블 기준 0.876을 기록했으며, 튜닝된 XGBoost와 CatBoost를 초월했다.
- 이전의 검색 증강 표준 모델보다 훨씬 더 효율적이며, 비용이 많이 드는 어텐션 또는 메모리 모듈을 피한 간결한 아키텍처를 채택했다.
- 제거 실험 결과, 어텐션 유사 검색 메커니즘이 성능 향상에 핵심적임을 확인했으며, 유사도 계산 및 이웃 집계의 적절한 설계가 일관된 성능 향상을 이끌었다.
더 나은 연구,지금 바로 시작하세요
논문 읽기부터 검토까지, 연구 시간을 획기적으로 줄여보세요.
카드 등록 없음 · 무료 플랜 제공
이 리뷰는 AI가 만들고, 인간 에디터가 검토했습니다.