Skip to main content
QUICK REVIEW

[논문 리뷰] Sequential Aggregation and Rematerialization: Distributed Full-batch Training of Graph Neural Networks on Large Graphs

Hesham Mostafa|arXiv (Cornell University)|2021. 11. 11.
Advanced Graph Neural Networks인용 수 4
한 줄 요약

이 논문은 그래프 신경망(GNNs)의 분산형 전체 배치 학습을 위한 순차적 집계 및 재물리화(Sequential Aggregation and Rematerialization, SAR)를 소개한다. SAR는 역전파 동안 계산 그래프 요소를 순차적으로 재구성하고 해제하여 메모리 사용을 극적으로 줄인다. SAR는 오브젝트 기반 분산 전체 배치 GNN 학습을 가능하게 하여, 각 워커의 메모리가 워커 수에 따라 선형적으로 증가하도록 보장함으로써 ogbn-papers100M 및 ogbn-products와 같은 큰 그래프에서 학습이 가능해지며, 128개의 워커에서 최대 50%의 메모리 절감을 달성하고 최적화된 어텐션 커널을 통해 빠른 성능 향상을 이룬다.

ABSTRACT

We present the Sequential Aggregation and Rematerialization (SAR) scheme for distributed full-batch training of Graph Neural Networks (GNNs) on large graphs. Large-scale training of GNNs has recently been dominated by sampling-based methods and methods based on non-learnable message passing. SAR on the other hand is a distributed technique that can train any GNN type directly on an entire large graph. The key innovation in SAR is the distributed sequential rematerialization scheme which sequentially re-constructs then frees pieces of the prohibitively large GNN computational graph during the backward pass. This results in excellent memory scaling behavior where the memory consumption per worker goes down linearly with the number of workers, even for densely connected graphs. Using SAR, we report the largest applications of full-batch GNN training to-date, and demonstrate large memory savings as the number of workers increases. We also present a general technique based on kernel fusion and attention-matrix rematerialization to optimize both the runtime and memory efficiency of attention-based models. We show that, coupled with SAR, our optimized attention kernels lead to significant speedups and memory savings in attention-based GNNs.We made the SAR GNN training library publicy available: \\url{https://github.com/IntelLabs/SAR}.

연구 동기 및 목표

  • 계산 그래프가 저장하기에 너무 커져서 처리가 불가능한 정도로 밀도가 높은 대규모 그래프에서 전체 배치 학습 시 메모리 병목 현상을 해결하기 위해.
  • 샘플링이나 비학습형 메시지 전파에 의존하지 않고도 어떤 GNN 아키텍처라도 확장 가능한 분산 전체 배치 학습을 가능하게 하기 위해.
  • 특히 어텐션 기반 모델에서 중간 텐서를 반복적으로 물리화하는 것을 방지함으로써 분산 학습에서의 통신 오버헤드를 최소화하기 위해.
  • 샘플링 기반 및 비학습형 메시지 전파 GNN 방법을 평가하기 위한 실용적이고 메모리 효율적인 기준을 제공하기 위해.
  • GAT와 같은 어텐션 기반 GNN의 성능을 최적화하기 위해 어텐션 계수를 저장하지 않고 실시간으로 계산함으로써 메모리 절약을 위해.

제안 방법

  • SAR는 입력 그래프를 N개의 워커에 분할하여 분산 도메인 병렬 학습 방식을 사용하며, 각 워커는 할당된 부분 그래프만 처리한다.
  • 전방 전파 동안 전체 GNN 계산 그래프를 물리적으로 저장하는 대신, SAR는 그래프 구축을 역전파까지 연기하여 조각별로 재구성한다.
  • 역전파 동안 SAR는 계산 그래프 요소를 순차적으로 재물리화하고 해제함으로써, 어떤 시점에도 최대 두 개의 파artition만 각 워커가 저장하도록 보장한다.
  • 이 방법은 워커당 메모리 스케일링을 O(2/N)로 구현하며, 그래프의 밀도와 무관하게 작동하므로 더 많은 워커를 추가함으로써 임의로 큰 그래프에서의 학습이 가능하다.
  • 어텐션 기반 모델의 경우, SAR는 커널 융합과 실시간 어텐션 계수 계산을 통합하여 큰 어텐션 행렬을 저장하지 않도록 한다.
  • 이 방법은 출력이 교차 워커 입력에 의존하는 모든 도메인 병렬 학습 설정, 예를 들어 공간 병렬 CNN에도 일반화 가능하다.

실험 결과

연구 질문

  • RQ1샘플링이나 비학습형 메시지 전파에 의존하지 않고도 전체 배치 GNN 학습을 대규모 그래프로 확장할 수 있는가?
  • RQ2분산 GNN 학습에서 메모리 사용을 어떻게 줄일 수 있을까? 이를 통해 메모리에 담을 수 없는 정도로 큰 그래프에서의 학습을 가능하게 할 수 있는가?
  • RQ3분산 환경에서 역전파 중 계산 그래프를 재물리화할 경우 통신 및 메모리 오버헤드는 얼마나 되는가?
  • RQ4GAT 모델에서 비용이 많이 드는 중간 텐서인 어텐션 계수를 재물리화하지 않도록 최적화할 수 있는가?
  • RQ5대규모 벤치마크에서 SAR는 샘플링 기반 및 비학습형 메시지 전파 GNN 방법과 비교해 메모리 효율성과 학습 속도 면에서 어떻게 성능을 내는가?

주요 결과

  • ogbn-papers100M에서 128개의 워커로 GraphSage 모델을 학습할 경우, SAR는 워커당 피크 메모리 소비를 최대 50%까지 줄였으며, 메모리 스케일링은 2/N로 선형적으로 유지된다.
  • ogbn-products에서 최적화된 어텐션 커널을 사용할 경우, DGL의 GAT 구현보다 SAR가 2.5배 빠른 성능을 보였다. 이는 메모리 압박 감소와 실시간 계수 계산 덕분이었다.
  • GraphSage의 경우, SAR는 통신량이 적은 도메인 병렬 학습과 유사한 런타임 성능을 달성하면서도 128개의 워커에서 메모리 사용량을 절반으로 줄였다.
  • SAR와 함께 사용된 최적화된 어텐션 커널(Free Attention Kernel, FAK)은 어텐션 계수를 저장하지 않음으로써 메모리 사용량을 줄였으며, 역전파 성능에 영향을 주지 않으면서도 전방 전파 속도를 향상시켰다.
  • SAR는 ogbn-papers100M 및 ogbn-products에서 현재까지 보고된 바 중 가장 큰 전체 배치 GNN 학습을 가능하게 하여, 1억 개 이상의 노드를 가진 그래프에서의 실행 가능성과 타당성을 입증했다.
  • 다양한 GNN 유형에 대해 통신을 피하는 방식이므로, 메모리 절감 효과가 런타임 오버헤드 없이 실현 가능하다.

더 나은 연구,지금 바로 시작하세요

논문 읽기부터 검토까지, 연구 시간을 획기적으로 줄여보세요.

카드 등록 없음 · 무료 플랜 제공

이 리뷰는 AI가 만들고, 인간 에디터가 검토했습니다.