연재 중집필 중인 책입니다. 아직 본문이 비어 있거나 채워지는 중인 장이 있습니다.

제 13 장

어텐션 커널

지금까지의 계산은 대부분 선형 층의 가중치를 기준으로 했다. 어텐션은 가중치가 없지만 문맥이 길어질수록 계산량과 메모리 접근이 길이에 따라 빠르게 늘어나 프리필과 디코드 모두에서 비중이 커진다. 서빙 엔진이 어떤 어텐션 커널을 쓰느냐는 긴 문맥 성능을 직접 좌우한다.

표준 어텐션의 메모리 왕복

길이 N의 시퀀스에서 헤드 하나의 어텐션은 세 단계로 계산한다. N은 문맥 길이, d는 헤드 차원이다.

S = Q × Kᵀ          (N × N 점수 행렬)
P = softmax(S)      (행마다 정규화)
O = P × V           (N × d 출력)

그대로 구현하면 S와 P를 HBM에 쓰고 다시 읽는다. N이 8,192면 점수 행렬 하나가 원소 6,700만 개로, BF16 기준 헤드 하나에 약 134MB다. 각 단계는 행렬 곱이나 원소별 연산이라 연산 자체는 빠르지만 이 큰 행렬을 HBM과 주고받는 시간이 전체를 지배한다. 연산 강도가 낮은 소프트맥스와 마스킹이 특히 그렇다.

GPU 안에는 HBM보다 훨씬 빠르지만 작은 온칩 SRAM(공유 메모리)이 있다. 점수 행렬 전체는 SRAM에 들어가지 않지만 작은 조각은 들어간다. 문제는 소프트맥스가 행 전체의 최댓값과 합을 알아야 계산된다는 점이다.

FlashAttention

FlashAttention(NeurIPS 2022)은 타일링으로 HBM과 SRAM 사이의 읽기·쓰기 횟수를 줄이는 IO 인지(IO-aware) 어텐션이다. Q, K, V를 블록으로 나눠 SRAM에 올리고, 블록 하나에서 점수와 소프트맥스와 출력까지 한 번에 계산한다. 점수 행렬은 HBM에 한 번도 쓰지 않는다.

행 전체를 보지 않고 소프트맥스를 계산하는 방법이 온라인 소프트맥스다. K 블록을 하나씩 처리할 때마다 지금까지의 행별 최댓값과 지수 합을 갱신하고, 최댓값이 바뀌면 이미 누적한 출력에 보정 계수를 곱한다. 마지막 블록까지 처리하면 전체 행을 한꺼번에 본 것과 같은 결과가 나온다. 근사가 아니라 정확한 어텐션이다.

표준 구현FlashAttention
HBM에 쓰는 중간 결과N × N 점수·확률 행렬없음 (행별 통계 N개만)
추가 메모리N²에 비례N에 비례
병목HBM 왕복행렬 곱 연산

후속판은 같은 원리 위에서 GPU를 더 잘 채운다. FlashAttention-2는 작업 분할과 병렬화를 고쳤고, FlashAttention-3(NeurIPS 2024)는 H100의 비동기 실행과 FP8을 써서 FP16 기준 최대 740 TFLOPS, H100 이론 성능의 75%에 도달했다고 보고한다. 같은 논문은 FlashAttention-2가 H100에서 35%에 그쳤다고 적는다. 커널 하나가 같은 장비의 실효 성능을 두 배 넘게 바꾼다.

프리필과 디코드의 어텐션

어텐션도 프리필에서는 연산에, 디코드에서는 대역폭에 묶인다. 선형 층에서 본 것과 같은 구도다.

프리필 어텐션은 쿼리 N개가 키 N개를 본다. 연산이 N²에 비례하고 같은 K, V 블록을 여러 쿼리 블록이 재사용하므로 연산 강도가 높다. FlashAttention 계열 커널로 연산 성능 상한에 가깝게 돈다. 프롬프트가 아주 길면 프리필에서 어텐션에 드는 시간이 선형 층보다 길어진다.

디코드 어텐션은 쿼리 1개가 그 요청의 KV 캐시 전체를 본다. KV 캐시를 한 번 읽어 쿼리 하나에만 쓰므로 연산 강도가 1 안팎이고, 시간은 KV 캐시를 읽는 대역폭이 정한다. GQA 모델에서는 쿼리 헤드 여러 개가 같은 KV 헤드를 공유하므로, 커널이 KV를 한 번 읽어 쿼리 헤드 여러 개에 함께 쓰면 연산 강도가 그 배수만큼 오른다.

디코드 어텐션에는 병렬성 문제도 있다. 배치가 작고 문맥이 길면 쿼리가 몇 개뿐이라 GPU의 SM을 다 채울 만큼 작업이 쪼개지지 않는다. 그래서 디코드용 커널은 KV 길이 방향으로 작업을 나눠 여러 SM이 한 요청의 KV 캐시 조각을 동시에 읽게 하고, 조각별 결과를 온라인 소프트맥스와 같은 방식으로 합친다.

페이지 단위 KV와 커널

블록 단위로 KV 캐시를 관리하면 한 요청의 KV가 메모리 여기저기 흩어진다. 어텐션 커널은 블록 테이블을 받아 블록 번호를 실제 주소로 바꿔 가며 읽는다. 블록 크기가 너무 작으면 한 번에 연속으로 읽는 양이 줄어 대역폭을 덜 쓰고, 블록 테이블 조회가 늘어난다. 블록 크기를 커널이 효율적으로 읽는 단위에 맞추는 이유다.

연속 배치에서는 한 스텝에 프리필 조각과 디코드 요청이 섞이고 요청마다 문맥 길이도 다르다. 엔진은 이런 가변 길이 배치를 커널 한 번으로 처리하는 어텐션 백엔드를 쓴다. 같은 엔진도 GPU 세대, KV 형식(BF16, FP8), 헤드 차원에 따라 다른 백엔드를 고르므로, 성능을 비교할 때는 엔진 로그에서 어떤 어텐션 백엔드가 선택됐는지 함께 기록한다.

고를 때 볼 것

운영자가 어텐션 커널을 직접 작성할 일은 드물다. 대신 엔진과 버전, 설정을 고르면서 다음을 확인한다.

  • 장비 세대에 맞는 커널이 쓰이는가. H100에서 이전 세대용 커널이 선택되면 FlashAttention-3 논문이 보인 것처럼 같은 장비에서 실효 성능이 크게 떨어진다
  • KV 형식과 커널이 맞는가. FP8 KV 캐시를 켰을 때 그 형식을 직접 읽는 커널이 있는지, 아니면 풀어서 읽느라 느려지는지 확인한다
  • 긴 문맥 부하로 잰다. 짧은 프롬프트로만 재면 어텐션 비중이 작아 커널 사이의 차이를 볼 수 없다. 서비스의 최대 문맥 근처까지 부하 테스트에 넣는다

디코딩 가속13 / 18