LabHub
배우기 러닝패스 코스

Transformer — アテンションを手で計算する

キー・バリューヘッドだけ減らしてみる

LabHub 에서 이어서 보기

한국어 원문으로 표시합니다.

목표

질의 헤드 수는 그대로 두고 키·값 헤드 수만 줄이는 구조를 표준 라이브러리만으로 만든다. 질의 헤드를 이어진 덩어리로 무리 짓고, 무리 안의 키·값 헤드를 평균으로 접고, 그것을 질의 헤드 수만큼 되풀이해 편 뒤 어텐션을 돌린다. 무리 수 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 원논문이 이미 학습된 멀티헤드 체크포인트를 옮길 때 쓰는 방식과 같지만, 원논문은 평균 뒤에 추가 학습을 한다. 여기서 나오는 출력 차이는 품질이 그만큼 나빠진다는 뜻이 아니라 같은 가중치로 키·값만 묶으면 답이 달라진다는 사실을 보여 주는 값이다. 채점기는 여러분이 적어 둔 설명을 믿지 않는다. 여러분의 모듈을 실제로 불러 매번 다른 헤드 수와 무리 수로 함수를 두드려 보고, 채점기가 따로 계산한 값과 대조한다.

단계

  1. /root/work/tf-gqa/gqa.pymake_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.pymake_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_membersgroup_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 이고 값은 정수입니다. scoresweightedh * n * n * head_dim, 투영은 n * d_model * (헤드 수 * head_dim) 꼴이며 total 은 나머지 여섯의 합입니다. 고르게 나눌 수 없으면 ValueError 입니다.

어느 항에 g 가 들어가고 어느 항에 안 들어가는지가 이 단계의 전부입니다. 키·값을 편 뒤에는 어텐션이 질의 헤드 수만큼 돌므로 scoresweighted 에는 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.jsonsample_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_multsscores + weighted 이고 무리 수가 바뀌어도 같은 값이라 core_same_for_all_g 가 참이 됩니다. md 에는 캐시 배수와 곱셈 배수를 나란히 적어 두세요 — 한쪽은 8배까지 줄고 다른 쪽은 거의 그대로라는 것이 이 실습의 결론입니다.