LabHub
배우기 러닝패스 코스

트랜스포머 — 어텐션을 손으로 계산한다 · MHA·MQA·GQA · 실습

키·값 헤드만 줄여 본다

LabHub 에서 이어서 보기

목표

질의 헤드 수는 그대로 두고 키·값 헤드 수만 줄이는 구조를 표준 라이브러리만으로 만든다. 질의 헤드를 이어진 덩어리로 무리 짓고, 무리 안의 키·값 헤드를 평균으로 접고, 그것을 질의 헤드 수만큼 되풀이해 편 뒤 어텐션을 돌린다. 무리 수 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.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 에 결과를 기록하세요.

참고

단계 8개

  1. 질의 헤드를 무리로 나눈다
  2. 무리 안의 키·값을 하나로 접는다
  3. 다시 질의 헤드 수만큼 편다
  4. 한 헤드짜리 어텐션
  5. 세 방식을 같은 가중치로
  6. 캐시 원소 수를 센다
  7. 곱셈 횟수를 센다
  8. 무엇이 줄고 무엇이 안 주는지 기록한다