디지털·가전제품
헤라아레스
분산 거대 모델 학습 시 텐서 파이프라인
안녕하세요! GPU 메모리에 대해 질문드려요
분산 거대 모델 학습 시 텐서 파이프라인 병렬화가 GPU 메모리 한계를 극복하는 기술적 원리는 무엇인가요?
1개의 답변이 있어요!
극복한다는 말 때문에 메모리를 마법처럼 늘려주는 기술로 오해하기 쉬운데, 실제로는 정반대입니다. 총 메모리는 그대로 두고 모델을 쪼개서 여러 장에 나눠 담는 거예요. 한 장이 감당해야 할 몫을 줄이는 것이지 없던 공간이 생기는 게 아닙니다.
먼저 왜 안 들어가는지부터 보셔야 합니다. 메모리를 먹는 게 파라미터만이 아니거든요. 학습할 때는 파라미터에 더해 기울기와 옵티마이저가 들고 있는 값, 그리고 역전파에 쓰려고 저장해 둔 중간 활성값까지 얹힙니다. 다 합치면 파라미터 용량의 몇 배가 되어서 카드 한 장에 안 들어가는 겁니다.
텐서 병렬은 층 하나를 가로로 자릅니다. 트랜스포머 안의 큰 행렬 곱을 열 단위로 나눠서 각 카드가 자기 몫만 계산한 뒤 결과를 합치는 방식이에요. 층 하나 자체가 너무 커서 안 들어갈 때 쓰는 방법입니다. 대신 층을 지날 때마다 카드끼리 결과를 주고받아야 해서 통신이 아주 잦습니다. 그래서 한 서버 안에서 고속으로 연결된 카드들끼리 묶는 게 보통이에요.
파이프라인 병렬은 반대로 층을 세로로 자릅니다. 앞쪽 열 개 층은 일번 카드, 다음 열 개는 이번 카드 이런 식이죠. 구간 경계에서 중간 결과만 넘기면 되니까 통신량이 훨씬 적어서 서버와 서버 사이에 씁니다. 문제는 일번이 끝나야 이번이 시작하니 뒤쪽 카드들이 손을 놓고 기다린다는 점인데, 이 노는 시간을 버블이라고 부릅니다.
버블은 배치를 잘게 쪼개서 줄입니다. 한 덩어리를 작은 조각들로 나눠서 첫 조각이 이번 카드로 넘어가는 순간 일번은 곧바로 다음 조각을 시작하는 거예요. 공장 컨베이어와 똑같습니다. 조각을 잘게 할수록 노는 시간이 줄지만 너무 잘게 나누면 한 번에 처리하는 효율이 떨어져서 적당한 지점을 찾아야 합니다.
실제 대규모 학습은 이 둘에 데이터 병렬까지 세 가지를 겹쳐서 씁니다. 여기에 옵티마이저 값을 카드들이 나눠 갖는 방식과, 중간 활성값을 저장하지 않고 역전파 때 다시 계산하는 방법을 더해 메모리를 더 짜냅니다. 결국 메모리를 아끼는 만큼 통신이나 계산이 늘어나는 구조라, 어디까지 감수할지를 정하는 저울질 문제에 가깝습니다.
채택된 답변