只把键值头减下来
한국어 원문으로 표시합니다.
목표
질의 헤드 수는 그대로 두고 키·값 헤드 수만 줄이는 구조를 표준 라이브러리만으로 만든다. 질의 헤드를 이어진 덩어리로 무리 짓고, 무리 안의 키·값 헤드를 평균으로 접고, 그것을 질의 헤드 수만큼 되풀이해 편 뒤 어텐션을 돌린다. 무리 수 g 를 바꾸는 것만으로 MHA(g=h)·GQA(1같은 가중치로 돌려 출력 차이를 재고, 마지막에 KV 캐시 원소 수와 곱셈 횟수를 각각 세어 무엇이 줄고 무엇이 안 주는지 확인한다.
왜 중요한가
요즘 모델 설정에는 질의 헤드 수와 키·값 헤드 수가 따로 적혀 있다. 두 값이 같으면 MHA, 키·값 쪽이 1 이면 MQA, 그 사이면 GQA 다. 세 이름을 외우는 것보다 무엇이 줄어드는가를 아는 것이 중요하다.
줄어드는 것은 KV 캐시 메모리다. 캐시 원소 수는 2 x 층 x 키·값 헤드 수 x 길이 x 헤드 차원 이고 이 식에 질의 헤드 수가 없다. 반면 어텐션 본체의 곱셈 횟수는 h x n x n x 헤드차원 이라 무리 수 g 가 아예 들어가지 않는다. 접은 키·값을 질의 헤드 수만큼 되풀이해 펴 놓고 평소대로 돌기 때문이다. 이 두 식을 직접 세어 보면 "GQA 로 바꿨는데 왜 프리필이 그대로냐" 는 물음에 답할 수 있다.
이 실습은 시간을 재지 않는다. 파드에 GPU 가 없고 CPU 도 다른 작업과 나눠 쓰므로 여기서 잰 속도는 아무것도 말해 주지 않는다. 판정은 전부 원소 수·곱셈 횟수 같은 정수 세기와 허용 오차를 둔 수치 대조다.
무리 안의 키·값 헤드를 평균으로 합치는 것은 이 실습의 가정이다. GQA 원논문이 이미 학습된 멀티헤드 체크포인트를 옮길 때 쓰는 방식과 같지만, 원논문은 평균 뒤에 추가 학습을 한다. 여기서 나오는 출력 차이는 품질이 그만큼 나빠진다는 뜻이 아니라 같은 가중치로 키·값만 묶으면 답이 달라진다는 사실을 보여 주는 값이다.
채점기는 여러분이 적어 둔 설명을 믿지 않는다. 여러분의 모듈을 실제로 불러 매번 다른 헤드 수와 무리 수로 함수를 두드려 보고, 채점기가 따로 계산한 값과 대조한다.
단계
- /root/work/tf-gqa/gqa.py 에
make_heads(count, rows, dim, seed)와group_of(head, h, g)·group_members(h, g)를 만드세요. 표본은 같은 인자면 늘 같은 값이 나와야 하고, 무리는 이어진 덩어리로 나눕니다. fold_kv(heads, g)를 더해 무리 안의 키·값 헤드를 자리마다 평균해 g개로 접게 하세요.expand_kv(folded, h)를 더해 접은 헤드를 제자리에서 되풀이해 h개로 펴게 하세요.attend(q, k, v)를 더하세요. 점수를sqrt(헤드 차원)으로 나누고, 소프트맥스는 최댓값을 뺀 뒤 지수를 잡습니다.heads_out(q_heads, k_heads, v_heads, g)와max_gap(left, right)를 더해 세 방식을 같은 가중치로 돌리고 출력 차이를 재게 하세요.kv_cache_elems(layers, kv_heads, seq_len, head_dim)와kv_table(layers, h, groups, seq_len, head_dim)을 더해 캐시 원소 수를 세게 하세요.mults(n, d_model, h, g, head_dim)을 더해 한 층 한 번 통과의 곱셈 횟수를 항목별로 세게 하세요.- 표본과 모델 규모를 정해 두 표를 만들고 /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 원논문 · GQA 원논문 · Attention Is All You Need · PyTorch — MultiheadAttention
- 흔한 실수: 무리를 번갈아 배정하기, 접을 때 첫 헤드만 남기기, 펼 때 목록 전체를 이어 붙이기, 점수를 나누지 않기, 소프트맥스에서 최댓값을 안 빼기, 캐시 식에 질의 헤드 수를 넣기, 곱셈 식의
h를g로 바꿔 계산량도 줄어든다고 세기.
질의 헤드를 무리로 나눈다
/root/work/tf-gqa/gqa.py 에 make_heads(count, rows, dim, seed) 와 group_of(head, h, g)·group_members(h, g) 를 만드세요. 표본은 같은 인자면 늘 같은 값이 나와야 하고(random 모듈 금지, 절댓값 4 이하), 무리는 이어진 덩어리로 나눕니다. h=8·g=2 면 0,1,2,3 이 0번 무리입니다. 고르게 나눌 수 없으면 ValueError 를 내세요.
표본은 작은 선형 합동식 하나면 충분합니다 — 상태를 정수로 들고 state = (a * state + c) % m 을 되풀이하면서 state / m - 0.5 를 꺼내면 같은 seed 에서 같은 수열이 나옵니다. 무리 번호는 head // (h // g) 입니다. head % g 는 번갈아 배정이라 뒤에서 펼 때 어긋납니다. group_members 는 group_of 를 h번 불러 채우면 두 곳의 규칙이 갈릴 일이 없습니다.
무리 안의 키·값을 하나로 접는다
fold_kv(heads, g) 를 더하세요. 무리 안의 키·값 헤드를 자리마다 평균해 g개로 접습니다. 헤드 수가 g 로 나누어떨어지지 않으면 ValueError 입니다. g 가 헤드 수와 같으면 값이 그대로 나와야 합니다.
무리는 1단계와 같은 규칙, 즉 이어진 덩어리입니다. heads[start:start + size] 로 한 무리를 떼어 자리마다 더하고 무리 크기로 나누면 됩니다. 첫 헤드만 남기거나 무리를 무시하고 전부 평균하면 g=1 에서만 우연히 같아 보입니다. 평균은 이 실습의 가정입니다 — GQA 원논문이 체크포인트를 옮길 때 쓰는 방식과 같지만, 원논문은 그 뒤에 추가 학습을 합니다.
다시 질의 헤드 수만큼 편다
expand_kv(folded, h) 를 더하세요. 접은 헤드를 제자리에서 되풀이해 h개로 폅니다. [A, B] 를 4개로 펴면 [A, A, B, B] 입니다. 고르게 펼 수 없으면 ValueError 입니다.
바깥 반복은 접은 헤드, 안쪽 반복은 h // len(folded) 번입니다. 목록 전체를 곱해 이어 붙이면(folded * size) [A, B, A, B] 가 되어 질의 헤드 i 가 남의 무리 키·값을 보게 됩니다. 1단계의 group_of 와 짝이 맞는지 확인하세요 — 질의 헤드 i 가 보는 것은 folded[group_of(i, h, g)] 여야 합니다. 행은 새 목록으로 복사해 돌려주면 나중에 한 곳을 고쳐 여러 헤드가 함께 바뀌는 일이 없습니다.
한 헤드짜리 어텐션
attend(q, k, v) 를 더하세요. 점수는 sqrt(헤드 차원) 으로 나누고, 소프트맥스는 최댓값을 뺀 뒤 지수를 잡습니다. 돌려주는 값은 q 와 같은 모양의 표입니다. 마스크는 쓰지 않습니다.
질의 행마다 모든 키 행과 내적을 하고, 나누고, 소프트맥스를 거쳐 값 행들의 가중 평균을 냅니다. sqrt 로 안 나누면 차원이 커질수록 소프트맥스가 한 자리로 쏠립니다. 최댓값을 빼지 않으면 점수가 큰 표본에서 math.exp 가 넘쳐 OverflowError 가 납니다 — 채점기가 일부러 큰 값을 넣어 봅니다.
세 방식을 같은 가중치로
heads_out(q_heads, k_heads, v_heads, g) 와 max_gap(left, right) 를 더하세요. heads_out 은 키·값을 g개로 접었다 h개로 편 뒤 질의 헤드마다 attend 를 한 번씩 부릅니다. g = len(q_heads) 면 평범한 멀티헤드와 한 자리도 다르지 않아야 하고, g 를 줄이면 출력이 달라져야 합니다. max_gap 은 두 출력 사이의 가장 큰 절댓값 차이를 돌려줍니다.
세 줄이면 끝납니다 — 키를 접었다 펴고, 값을 접었다 펴고, 질의 헤드마다 attend. 질의 헤드는 건드리지 않는다는 점이 이 단계의 전부입니다. max_gap 은 헤드·행·열을 전부 훑어 가장 큰 차이 하나만 남깁니다. g 를 바꿔 가며 불러 보면 g=h 에서 차이가 정확히 0 이고 g 가 작아질수록 차이가 커지는 것을 볼 수 있습니다.
캐시 원소 수를 센다
kv_cache_elems(layers, kv_heads, seq_len, head_dim) 와 kv_table(layers, h, groups, seq_len, head_dim) 을 더하세요. 원소 수는 2 * layers * kv_heads * seq_len * head_dim 이고 질의 헤드 수는 들어가지 않습니다. kv_table 은 [(무리 수, 원소 수), ...] 를 돌려주고, 고르게 나눌 수 없는 무리 수가 있으면 ValueError 입니다.
앞의 2 는 K 와 V 두 벌입니다. 질의 헤드 수를 곱하고 싶어지는 자리가 있는데, 질의는 캐시에 남기지 않으므로 식에 없습니다. kv_table 에서 그 무리 수의 키·값 헤드 수가 곧 g 라는 점을 그대로 쓰면 됩니다. 바이트가 아니라 원소 수로 세는 이유는 자료형에 따라 한 원소가 2바이트도 4바이트도 되기 때문입니다.
곱셈 횟수를 센다
mults(n, d_model, h, g, head_dim) 을 더하세요. 열쇠는 proj_q·proj_k·proj_v·scores·weighted·proj_out·total 이고 값은 정수입니다. scores 와 weighted 는 h * n * n * head_dim, 투영은 n * d_model * (헤드 수 * head_dim) 꼴이며 total 은 나머지 여섯의 합입니다. 고르게 나눌 수 없으면 ValueError 입니다.
어느 항에 g 가 들어가고 어느 항에 안 들어가는지가 이 단계의 전부입니다. 키·값을 편 뒤에는 어텐션이 질의 헤드 수만큼 돌므로 scores 와 weighted 에는 g 가 없습니다. g 에 딸린 것은 K 투영과 V 투영 둘뿐입니다. 곱셈만 세고 덧셈과 소프트맥스의 지수는 세지 않습니다 — 보려는 것은 총량이 아니라 어느 항이 줄어드는가입니다.
무엇이 줄고 무엇이 안 주는지 기록한다
표본은 make_heads(8, 6, 4, 101)·make_heads(8, 6, 4, 202)·make_heads(8, 6, 4, 303) 을 Q·K·V 로, 무리 수는 [8, 4, 2, 1] 로 둡니다. 모델 규모는 layers = 32, seq_len = 4096, head_dim = 128, h = 8, d_model = 1024, n_tokens = 4096 입니다. /root/work/tf-gqa/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 를, /root/work/tf-gqa/gqa_report.md 에 ## 무엇을 쟀나 ## 무리를 줄이면 출력이 얼마나 달라지나 ## 메모리는 줄고 곱셈은 안 준다 ## 어디에 쓸 것인가 네 절로 쓰세요.
숫자는 손으로 적지 말고 여러분의 함수를 실제로 돌려 얻은 값으로 채우세요. diff_table 은 무리 수마다 max_gap(그 무리의 출력, g=h 의 출력) 이라 첫 칸이 0.0 입니다. kv_ratio·mult_ratio 는 MHA 를 1 로 둔 배수이므로 MHA 값 / 그 무리의 값 입니다. core_mults 는 scores + weighted 이고 무리 수가 바뀌어도 같은 값이라 core_same_for_all_g 가 참이 됩니다. md 에는 캐시 배수와 곱셈 배수를 나란히 적어 두세요 — 한쪽은 8배까지 줄고 다른 쪽은 거의 그대로라는 것이 이 실습의 결론입니다.