트랜스포머 — 어텐션을 손으로 계산한다 · MHA·MQA·GQA · 실습
키·값 헤드만 줄여 본다
목표
질의 헤드 수는 그대로 두고 키·값 헤드 수만 줄이는 구조를 표준 라이브러리만으로 만든다. 질의 헤드를 이어진 덩어리로 무리 짓고, 무리 안의 키·값 헤드를 평균으로 접고, 그것을 질의 헤드 수만큼 되풀이해 편 뒤 어텐션을 돌린다. 무리 수 g 를 바꾸는 것만으로 MHA(g=h)·GQA(1<g<h)·MQA(g=1) 세 방식을 같은 가중치로 돌려 출력 차이를 재고, 마지막에 KV 캐시 원소 수와 곱셈 횟수를 각각 세어 무엇이 줄고 무엇이 안 주는지 확인한다.
왜 중요한가
요즘 모델 설정에는 질의 헤드 수와 키·값 헤드 수가 따로 적혀 있다. 두 값이 같으면 MHA, 키·값 쪽이 1 이면 MQA, 그 사이면 GQA 다. 세 이름을 외우는 것보다 무엇이 줄어드는가를 아는 것이 중요하다.
줄어드는 것은 KV 캐시 메모리다. 캐시 원소 수는 2 x 층 x 키·값 헤드 수 x 길이 x 헤드 차원 이고 이 식에 질의 헤드 수가 없다. 반면 어텐션 본체의 곱셈 횟수는 h x n x n x 헤드차원 이라 무리 수 g 가 아예 들어가지 않는다. 접은 키·값을 질의 헤드 수만큼 되풀이해 펴 놓고 평소대로 돌기 때문이다. 이 두 식을 직접 세어 보면 "GQA 로 바꿨는데 왜 프리필이 그대로냐" 는 물음에 답할 수 있다.
이 실습은 시간을 재지 않는다. 파드에 GPU 가 없고 CPU 도 다른 작업과 나눠 쓰므로 여기서 잰 속도는 아무것도 말해 주지 않는다. 판정은 전부 원소 수·곱셈 횟수 같은 정수 세기와 허용 오차를 둔 수치 대조다.
무리 안의 키·값 헤드를 평균으로 합치는 것은 이 실습의 가정이다. GQA 원논문이 이미 학습된 멀티헤드 체크포인트를 옮길 때 쓰는 방식과 같지만, 원논문은 평균 뒤에 추가 학습을 한다. 여기서 나오는 출력 차이는 품질이 그만큼 나빠진다는 뜻이 아니라 같은 가중치로 키·값만 묶으면 답이 달라진다는 사실을 보여 주는 값이다.
채점기는 여러분이 적어 둔 설명을 믿지 않는다. 여러분의 모듈을 실제로 불러 매번 다른 헤드 수와 무리 수로 함수를 두드려 보고, 채점기가 따로 계산한 값과 대조한다.
단계
1. /root/work/tf-gqa/gqa.py 에 make_heads(count, rows, dim, seed) 와 group_of(head, h, g)·group_members(h, g) 를 만드세요. 표본은 같은 인자면 늘 같은 값이 나와야 하고, 무리는 이어진 덩어리로 나눕니다.
2. fold_kv(heads, g) 를 더해 무리 안의 키·값 헤드를 자리마다 평균해 g개로 접게 하세요.
3. expand_kv(folded, h) 를 더해 접은 헤드를 제자리에서 되풀이해 h개로 펴게 하세요.
4. attend(q, k, v) 를 더하세요. 점수를 sqrt(헤드 차원) 으로 나누고, 소프트맥스는 최댓값을 뺀 뒤 지수를 잡습니다.
5. heads_out(q_heads, k_heads, v_heads, g) 와 max_gap(left, right) 를 더해 세 방식을 같은 가중치로 돌리고 출력 차이를 재게 하세요.
6. kv_cache_elems(layers, kv_heads, seq_len, head_dim) 와 kv_table(layers, h, groups, seq_len, head_dim) 을 더해 캐시 원소 수를 세게 하세요.
7. mults(n, d_model, h, g, head_dim) 을 더해 한 층 한 번 통과의 곱셈 횟수를 항목별로 세게 하세요.
8. 표본과 모델 규모를 정해 두 표를 만들고 /root/work/tf-gqa/gqa_report.json 과 /root/work/tf-gqa/gqa_report.md 에 결과를 기록하세요.
참고
- 실행 계약: 채점기는
/root/work/tf-gqa/gqa.py를 파이썬 모듈로 불러make_heads·group_of·group_members·fold_kv·expand_kv·attend·heads_out·max_gap·kv_cache_elems·kv_table·mults를 직접 씁니다. 스크립트로 실행하지 않으므로if __name__ == "__main__"은 없어도 됩니다. make_heads(count, rows, dim, seed)는 헤드 count개를 돌려주고 각 헤드는 rows개의 행, 한 행은 dim개의 실수입니다. 같은 인자면 늘 같은 값이어야 하고(random모듈 금지) 다른 seed 면 다른 값이어야 합니다. 값은 절댓값 4 이하로 두고, 모든 값이 같으면 안 됩니다.group_of(head, h, g)는 이어진 덩어리로 배정합니다. h=8, g=2 면 0,1,2,3 이 0번 무리이고 4,5,6,7 이 1번 무리입니다.h % g != 0이거나 h·g 가 1 보다 작거나 head 가 범위를 벗어나면ValueError를 냅니다.group_members(h, g)는 길이 g 의 목록이고 각 칸은 그 무리에 든 질의 헤드 번호 목록입니다.fold_kv(heads, g)는 무리 안의 헤드를 자리마다 평균해 g개로 접습니다. g 가 헤드 수와 같으면 값이 그대로 나옵니다(무리마다 헤드가 하나라 평균이 자기 자신입니다). 고르게 나눌 수 없으면ValueError입니다.expand_kv(folded, h)는 각 헤드를 제자리에서 되풀이합니다.[A, B]를 4개로 펴면[A, A, B, B]이지[A, B, A, B]가 아닙니다. 고르게 펼 수 없으면ValueError입니다.attend(q, k, v)는 q 와 같은 모양의 표를 돌려줍니다. 마스크는 쓰지 않습니다 — 여기서 보려는 것은 키·값 헤드 수의 효과뿐이라 인과 마스크를 빼고 견줍니다.heads_out(...)는fold_kv로 접고expand_kv로 편 뒤 질의 헤드마다attend를 한 번씩 부릅니다.g = len(q_heads)로 부르면 접기·펴기가 값을 바꾸지 않으므로 평범한 멀티헤드와 같은 결과가 나와야 합니다.max_gap(left, right)는 두 출력 사이의 가장 큰 절댓값 차이 하나를 돌려줍니다.kv_cache_elems는2 * layers * kv_heads * seq_len * head_dim입니다. 질의 헤드 수는 들어가지 않습니다.kv_table은[(무리 수, 원소 수), ...]입니다.mults(n, d_model, h, g, head_dim)의 열쇠는proj_q·proj_k·proj_v·scores·weighted·proj_out·total일곱 개이고 값은 전부 정수입니다.total은 나머지 여섯의 합입니다. 곱셈만 세고 덧셈과 지수는 세지 않습니다.- 8단계 보고서의 값은 이렇게 정합니다. 출력 비교는
sample_h = 8,sample_n = 6,sample_head_dim = 4에make_heads(8, 6, 4, 101)·make_heads(8, 6, 4, 202)·make_heads(8, 6, 4, 303)을 각각 Q·K·V 로 씁니다. 무리 수는groups = [8, 4, 2, 1]입니다. - 모델 규모는
layers = 32,seq_len = 4096,head_dim = 128,h = 8,d_model = 1024,n_tokens = 4096으로 둡니다. 실제 모델을 재어 온 값이 아니라 표를 만들기 위한 예시 규모입니다. gqa_report.json에 넣을 열쇠:sample_h·sample_n·sample_head_dim·groups·diff_table·layers·seq_len·head_dim·h·d_model·n_tokens·kv_table·kv_ratio·mult_table·mult_ratio·core_mults·core_same_for_all_g.diff_table은[[무리 수, MHA 와의 최대 차이], ...],kv_ratio·mult_ratio는 MHA 를 1 로 둔 배수(MHA 값 / 그 무리의 값),core_mults는scores + weighted입니다.gqa_report.md는## 무엇을 쟀나## 무리를 줄이면 출력이 얼마나 달라지나## 메모리는 줄고 곱셈은 안 준다## 어디에 쓸 것인가네 절로 씁니다.- 이 파드에는 인터넷도 GPU 도 없습니다.
pip install은 되지 않고 numpy 는/opt/onnx-lab/bin/python안에만 있어 시스템 파이썬에서는import numpy가 되지 않습니다.math만으로 충분합니다. - 공식 문서: [MQA 원논문](https://arxiv.org/abs/1911.02150) · [GQA 원논문](https://arxiv.org/abs/2305.13245) · [Attention Is All You Need](https://arxiv.org/abs/1706.03762) · [PyTorch — MultiheadAttention](https://docs.pytorch.org/docs/stable/generated/torch.nn.MultiheadAttention.html)
- 흔한 실수: 무리를 번갈아 배정하기, 접을 때 첫 헤드만 남기기, 펼 때 목록 전체를 이어 붙이기, 점수를 나누지 않기, 소프트맥스에서 최댓값을 안 빼기, 캐시 식에 질의 헤드 수를 넣기, 곱셈 식의
h를g로 바꿔 계산량도 줄어든다고 세기.
단계 8개
- 질의 헤드를 무리로 나눈다
- 무리 안의 키·값을 하나로 접는다
- 다시 질의 헤드 수만큼 편다
- 한 헤드짜리 어텐션
- 세 방식을 같은 가중치로
- 캐시 원소 수를 센다
- 곱셈 횟수를 센다
- 무엇이 줄고 무엇이 안 주는지 기록한다