[논문 리뷰] HelixFold: An Efficient Implementation of AlphaFold2 using PaddlePaddle
HelixFold는 PaddlePaddle에서 AlphaFold2를 엔드투엔드로 구현하고 Branch Parallelism 및 메모리 최적화를 통해 학습을 더 빠르게 하면서 CASP14와 CAMEO에서 AlphaFold2와 비교 가능한 정확도를 유지합니다.
Accurate protein structure prediction can significantly accelerate the development of life science. The accuracy of AlphaFold2, a frontier end-to-end structure prediction system, is already close to that of the experimental determination techniques. Due to the complex model architecture and large memory consumption, it requires lots of computational resources and time to implement the training and inference of AlphaFold2 from scratch. The cost of running the original AlphaFold2 is expensive for most individuals and institutions. Therefore, reducing this cost could accelerate the development of life science. We implement AlphaFold2 using PaddlePaddle, namely HelixFold, to improve training and inference speed and reduce memory consumption. The performance is improved by operator fusion, tensor fusion, and hybrid parallelism computation, while the memory is optimized through Recompute, BFloat16, and memory read/write in-place. Compared with the original AlphaFold2 (implemented with Jax) and OpenFold (implemented with PyTorch), HelixFold needs only 7.5 days to complete the full end-to-end training and only 5.3 days when using hybrid parallelism, while both AlphaFold2 and OpenFold take about 11 days. HelixFold saves 1x training time. We verified that HelixFold's accuracy could be on par with AlphaFold2 on the CASP14 and CAMEO datasets. HelixFold's code is available on GitHub for free download: https://github.com/PaddlePaddle/PaddleHelix/tree/dev/apps/protein_folding/helixfold, and we also provide stable web services on https://paddlehelix.baidu.com/app/drug/protein/forecast.
연구 동기 및 목표
- AlphaFold2의 정확도를 해치지 않으면서 학습 시간과 메모리 사용량을 줄이는 것을 목표로 한다
- PaddlePaddle(HelixFold)에서 AlphaFold2의 완전한 학습 및 추론 파이프라인을 개발한다
- 효율적인 엔드-투-엔드 학습을 가능하게 하는 새로운 병렬성 및 메모리 최적화 기법을 도입한다
제안 방법
- 다수의 연산자 및 텐서를 융합하여 CPU 스케줄링 오버헤드를 줄인다(텐서 융합 및 융합 게이트드 셀프 어텐션)
- Evoformer의 MSA 및 Pair 스택을 디바이스 간에 병렬화하기 위한 Branch Parallelism(BP) 도입
- 데이터 병렬성과 메모리 효율성을 높이기 위한 하이브리드 병렬성(BP-DAP-DP) 채택
- 재계산(Recomputation), BFloat16 활성화, 제로TI/메모리 읽기/쓰기, Subbatch/DAP 기법 등 메모리 최적화 적용
- 커널 런칭 수와 메모리 단편화를 줄이기 위해 매개변수, 그래디언트, 옵티마이저 상태를 융합
- Very long sequence의 피크 메모리 관리에 Subbatch 및 DAP 활용
실험 결과
연구 질문
- RQ1HelixFold가 CASP14 및 CAMEO에서 학습 시간을 줄이면서 AlphaFold2의 정확도에 맞출 수 있는가?
- RQ2융합, 브랜치 병렬성, 메모리 최적화 기법이 학습 처리량과 메모리 풋프린트에 어떤 영향을 미치는가?
- RQ3BP-DAP-DP 하이브리드 병렬성이 Evoformer의 기존 병렬화 전략 대비 어떤 이점을 제공하는가?
- RQ4PaddlePaddle 기반 HelixFold를 사용할 때 원래의 AlphaFold2 및 OpenFold에 비해 엔드-투-엔드 학습 시간 및 자원 비용이 얼마나 감소하는가?
주요 결과
- HelixFold는 CASP14에서 AlphaFold2와의 정확도에서 경쟁력 있는 성능을 달성합니다(TM-score 0.8771 vs 0.8772) 및 CAMEO에서의 성능도 유사합니다(TM-score 0.8885 vs 0.8862)
- 엔드-투-엔드 학습 시간이 7.5일로 단축되었으며(하이브리드 병렬성으로 5.3일), AlphaFold2 및 OpenFold의 약 11일과 비교
- 초기 학습에서 처리량이 향상되었고: HelixFold가 AlphaFold2 대비 처리량이 +28.42% ~ +52.13% 증가 및 OpenFold보다 높았으며, 미세조정에서 약 +84.49%로 AlphaFold2를 능가
- BP-DAP-DP 하이브리드 병렬성이 DAP 단독보다 효율이 높고 통신 오버헤드가 더 작다
- Recompute, BFloat16, in-place read/write 및 텐서 융합 등의 메모리 최적화가 피크 메모리 및 커널 런칭 수를 크게 줄인다
- 융합된 Gated Self-Attention이 정확도는 유지하면서 CPU/GPU 오버헤드를 감소시킨다
더 나은 연구,지금 바로 시작하세요
논문 읽기부터 검토까지, 연구 시간을 획기적으로 줄여보세요.
카드 등록 없음 · 무료 플랜 제공
이 리뷰는 AI가 만들고, 인간 에디터가 검토했습니다.