저자들은 이 책에서 트랜스포머 모델을 대규모 하드웨어에 얹을 때 무엇이 속도를 정하는지를 원리로 풀어낸다. 강한 스케일링(칩을 늘린 만큼 처리량이 비례해서 늘어나는 것)이 무너지는 자리는 칩 사이 통신이 계산 시간보다 길어지는 순간이라고 짚는다. 500조 연산/초로 표시된 가속기가 메모리 이동에 발목 잡히면 실제로는 표시치의 10분의 1만 낼 수 있다는 예시를 든다. 4부 12장으로 짜여, 1~3장은 루프라인·TPU·샤딩 표기를, 4~8장은 트랜스포머의 파라미터·FLOPs 계산과 데이터·텐서·파이프라인·엑스퍼트 네 가지 병렬화, LLaMA 3 학습·서빙 실습을 다룬다. 9~10장은 JAX 프로파일링과 프로그래밍을, 12장은 GPU를 새로 다룬다. 저자들은 제임스 브래드버리와 블레이크 헥트먼의 아이디어를 많이 빌렸다고 밝힌다.
한줄 코멘트. 이 책이 파는 것은 신비가 아니라 계산이다. 저자들은 모델을 키우는 일이 통신과 메모리라는 두 병목의 계산으로 환원된다고 말하고, 그 계산을 익히면 실제 하드웨어를 돌려 보지 않고도 어느 병렬화가 맞을지 가늠할 수 있다고 본다. 이 장은 12장짜리 책의 지도이고, 구체적인 숫자는 뒤 장에서 나온다.
3~4년 전만 해도 대다수 머신러닝 연구자는 이 책에 나오는 내용을 몰라도 됐다고 저자들은 말한다. 지금은 "작은" 모델조차 하드웨어 한계에 바짝 붙어 돌아가서, 새로운 연구를 하려면 규모에서의 효율을 같이 생각해야 한다. 저자들이 짚는 과거 사례는 알렉스 크리제프스키다. CNN(합성곱 신경망)을 빠르게 돌리려고 날것의 CUDA 코드를 직접 짰던 그의 작업은 몇 년 뒤 Theano·TensorFlow 같은 라이브러리가 대신 처리해 줬다. 저자들은 지금 책에 담은 내용도 몇 년 안에 그렇게 추상화될 수 있다고 인정하면서도, 스케일링 법칙이 모델을 계속 하드웨어의 한계선까지 밀어붙이는 한 최전선 연구는 대규모로 모델을 효율적으로 돌리는 법과 떼어 놓을 수 없을 거라고 본다. 저자들의 표현을 빌리면 "벤치마크에서 20% 이기더라도 루프라인 효율에서 20%를 깎아 먹으면 의미가 없다." 유망한 모델 구조가 실패하는 이유도 대개 둘 중 하나다. ① 규모에서 효율적으로 못 돌거나, ② 그렇게 돌아가게 만드는 작업에 아무도 공을 들이지 않아서다.
저자들이 "모델 스케일링"이라 부르는 목표는 단순하다. 학습이나 추론에 쓰는 칩 수를 늘릴 때 처리량도 그만큼 비례해서 늘리는 것, 이것이 "강한 스케일링(strong scaling)"이다. 칩을 더 붙이는 병렬화는 계산 시간을 줄여 주지만 그 대가로 칩 사이에 오가는 통신이 늘어난다. 통신에 걸리는 시간이 계산 시간을 넘어서면 "통신에 발목 잡힌(communication bound)" 상태가 되고, 그 순간부터는 칩을 더 붙여도 처리량이 비례해서 늘지 않는다. 계산 시간이 줄면 이번엔 칩 한 개 수준의 병목이 드러난다. 저자들은 500조 연산/초를 낸다고 표시된 TPU나 GPU라도 파라미터를 메모리에서 옮기는 데 발목 잡히면 표시치의 10분의 1만 낼 수 있다고 짚는다. 칩 하나가 처리하는 연산량과 메모리 대역폭, 전체 메모리 용량이 스케일링 이야기의 핵심에 있는 이유다. 이 병목이 어디서 나타날지 미리 알면 그것을 피하도록 모델을 설계하거나 다시 짤 수 있다는 것이 저자들의 논리다.
하드웨어를 설계하는 쪽은 반대편 문제를 짊어진다. 비용을 최소로 두면서 알고리즘이 딱 필요한 만큼의 연산·대역폭·메모리를 내주는 하드웨어를 만들어야 한다. 저자들은 이 코디자인(co-design, 하드웨어와 알고리즘을 함께 걸고 설계하는 문제)이 얼마나 위태로운지 짚는다. 실제 칩이 나오기까지 2~3년이 걸리는데, 그사이 알고리즘이 어떤 모습일지 미리 걸어야 한다. TPU는 이 내기에서 이긴 사례로 저자들이 꼽는 이야기다. 행렬곱은 메모리 바이트당 소화하는 FLOPs(부동소수점 연산 수)가 다른 어떤 연산보다도 많은(바이트당 N FLOPs) 독특한 알고리즘이고, 시스톨릭 배열(데이터가 격자 모양으로 늘어선 연산 유닛 사이를 리듬감 있게 흘러가며 계산되는 구조)로 지은 초기 TPU는 나온 시점의 GPU보다 달러당 성능에서 훨씬 앞섰다고 저자들은 말한다. TPU는 처음부터 머신러닝 작업에 맞춰 설계됐고, 텐서 코어를 얹은 GPU도 빠르게 같은 틈을 메워 가는 중이다. 반대로 신경망이 그때 뜨지 않았거나 TPU가 다루기 힘든 방향으로 근본적으로 바뀌었다면, GPU보다 유연성이 떨어지는 TPU에 건 그 내기는 값비싼 실패로 남았을 거라고 저자들은 짚는다.
12장을 4부로 어떻게 엮었나
책은 4부 12장으로 짜여 있다. 1부(1~3장)는 예비지식이다. 1장은 계산·통신·메모리 세 가지가 알고리즘 속도를 어떻게 가두는지 다루는 루프라인(roofline, 무엇이 속도를 묶는지 재는 분석) 분석이고, 2장은 TPU가 칩 하나로서 그리고 제한된 대역폭·지연시간을 가진 칩 사이 연결로 묶인 시스템으로서 어떻게 동작하는지, 3장은 여러 TPU에 흩어진 행렬을 어떻게 곱하는지를 샤딩(sharding, 행렬을 여러 칩에 쪼개 나눠 담는 것)으로 설명한다. 이 세 장에서 저자들은 행렬곱 하나가 계산에 발목 잡히는지 메모리·통신에 발목 잡히는지, TPU가 어떻게 학습 클러스터로 배선되고 부분마다 대역폭을 얼마나 갖는지, 여러 TPU에 흩어진 배열을 모으고 흩뿌리고 다시 나누는 데 걸리는 시간과 서로 다르게 흩어진 행렬을 효율적으로 곱하는 법을 답한다. 2부(4~8장)는 트랜스포머다. 4장은 순전파·역전파에 드는 FLOPs, 파라미터 수, KV 캐시(어텐션에 쓰는 키·값을 저장해 두는 캐시) 크기까지 트랜스포머 수학을 다루고, 이 계산으로 모델이 메모리를 얼마나 쓰는지, 계산과 통신에 시간을 얼마나 쓰는지, 어텐션이 피드포워드 블록에 비해 언제 중요해지는지를 알 수 있다. 5장과 7장, 학습과 추론 장이 이 책의 중심이라고 저자들은 밝힌다. 모델 크기와 칩 수가 주어졌을 때 어떻게 나눠야 강한 스케일링 영역에 머무는지를 묻는 물음이다. 6장과 8장은 이 개념들을 LLaMA 3라는 널리 쓰는 오픈소스 모델에 적용하는 실습이다. 3부(9~10장)는 9장이 JAX+XLA 스택과 JAX/텐서보드 프로파일러로 실제 문제를 디버깅하는 법을, 10장이 계산을 병렬화하는 JAX API를 예제로 다룬다. 4부(11~12장)는 마무리 장과, GPU가 어떻게 동작하고 어떻게 배선되며 루프라인이 TPU와 어떻게 다른지 새로 다루는 12장이다.
무엇을 쪼개고 무엇을 줄이나
모델을 여러 칩에 나눠 쓰는 데는 두 갈래 기법이 있다고 저자들은 짚는다. ① 계산을 쪼개는 병렬화 넷(데이터·텐서·파이프라인·엑스퍼트), ② 메모리 요구량 자체를 줄이는 기법 몇 가지다. 재계산(rematerialization, 중간 계산값을 저장하는 대신 필요할 때 다시 계산하는 것), 옵티마이저·모델 샤딩(ZeRO로 불리는 방식), 호스트 오프로드(파라미터를 호스트 메모리로 내보내는 것), 그레이디언트 누적이 후자에 들어간다. 5장과 7장에서 이 목록을 자세히 다루며, 주어진 칩 수와 모델 크기에서 어떤 조합을 골라야 강한 스케일링 영역에 머무는지를 저자들은 풀어낸다. 책을 처음부터 끝까지 순서대로 읽을 필요는 없다고 저자들은 밝힌다. 1~3장은 전제 지식과 뒤에서 쓸 표기를 세우는 자리라 이미 익숙하면 건너뛰어도 된다. 저자들은 책 끝에 제임스 브래드버리와 블레이크 헥트먼이 이 책에 담긴 여러 아이디어를 이끌어 냈다고 밝혀 둔다.