LabHub
배우기 러닝패스 코스

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

KV キャッシュが節約する乗算を数える

LabHub 에서 이어서 보기

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

목표

자기회귀 생성에서 KV 캐시가 무엇을 아끼는지 곱셈 횟수를 직접 세어 확인한다. 곱셈을 세는 곱셈 함수를 만들고, 캐시 없이 도는 판과 캐시로 도는 판을 각각 만들어 두 판의 출력이 같은 값인지 허용 오차로 확인한 뒤, 문맥 길이마다 곱셈 횟수를 세어 표로 만든다. 마지막에 캐시를 낮은 정밀도로 들고 있으면 그 같음이 깨지는 것을 재고, 자기가 정한 설정으로 캐시가 먹는 원소 수와 바이트를 계산해 기록으로 남긴다.

왜 중요한가

어텐션 한 번을 손으로 계산하는 것과, 그 어텐션을 생성 반복 안에서 부르는 것은 다른 이야기다. 길이 100짜리 답을 만들려면 모델을 100번 부르고, 캐시가 없으면 부를 때마다 앞 문맥 전체의 k·v 를 처음부터 다시 만든다. 첫 토큰의 k·v 를 백 번 만들고 아흔아홉 번 버린다는 뜻이다. 이 낭비는 코드를 읽어서는 보이지 않는다. 어텐션 함수 자체는 잘못된 곳이 하나도 없기 때문이다. 보이게 하려면 세어야 한다. 그리고 시간을 재서는 안 된다 — 시간은 기계와 부하에 따라 달라지지만 곱셈 횟수는 같은 입력에서 늘 같고, 길이에 대해 어떻게 늘어나는지를 그대로 보여 준다. 이 실습은 실제 모델을 부르지 않는다. 이 파드의 시스템 파이썬에는 numpy 가 없고(/opt/onnx-lab/bin/python 안에만 있다) torch 도 transformers 도 없다. 표준 라이브러리만으로 같은 구조를 만들고, 여기서 잰 숫자만 쓴다. 그래서 "실제 모델은 몇 GB 를 먹는다" 나 "몇 배 빨라진다" 같은 말은 여기서 하지 않는다. 채점기는 여러분이 적어 둔 설명을 믿지 않는다. 여러분의 모듈을 실제로 불러 매번 다른 크기의 입력으로 함수를 두드려 보고, 값과 곱셈 횟수를 따로 계산해 대조한다. 크기가 실행마다 바뀌므로 값을 외워 넣을 수 없다.

단계

  1. /root/work/tf-kv/kv.pyCONFIGreset_muls()·muls()·mul(a, b)·dot(u, v) 를 만드세요. mul 은 곱하면서 센 횟수를 하나 올리고, dot 은 그 mul 로만 곱합니다.
  2. project(x, W) 를 더해 토큰 하나를 가중치 한 벌에 통과시키게 하세요. W 는 줄마다 길이가 len(x) 인 행렬이고 돌려주는 목록의 길이는 len(W) 입니다.
  3. attend(q, K, V) 를 더해 질의 하나가 쌓인 K·V 전체를 보게 하세요. 점수를 √len(q) 로 나누고 소프트맥스를 씌운 뒤 V 의 가중 평균을 냅니다.
  4. step_nocache(xs, Wq, Wk, Wv) 를 더해 캐시 없이 출력 하나를 내게 하세요. 앞의 모든 토큰에 대해 q·k·v 를 다시 만들고 마지막 질의로 어텐션합니다.
  5. new_cache()·store(cache, k, v, digits=None)·step_cached(x, Wq, Wk, Wv, cache, digits=None) 를 더해 캐시로 같은 출력을 내게 하세요. 새 토큰 하나만 계산하고 K·V 에 한 줄 덧붙입니다.
  6. ATOL·RTOL·close(got, want)·compare(xs, Wq, Wk, Wv, digits=None) 를 더해 두 방식의 출력이 같은 값인지 견주게 하세요. 등호로 견주지 않습니다.
  7. mul_table(xs, Wq, Wk, Wv) 를 더해 문맥 길이 1 부터 len(xs) 까지 두 방식의 곱셈 횟수를 세게 하세요. 돌려주는 값은 (문맥길이, 캐시없음, 캐시있음) 짝의 목록입니다.
  8. 토큰 12개로 표를 만들고 캐시가 먹는 메모리를 계산해 /root/work/tf-kv/kv_report.json/root/work/tf-kv/kv_report.md 에 기록하세요.

참고

곱셈을 세는 곱셈

/root/work/tf-kv/kv.pyCONFIGreset_muls()·muls()·mul(a, b)·dot(u, v) 를 만드세요. CONFIGlayers·heads·head_dim·dtype_bytes 네 열쇠를 가진 딕셔너리이고 값은 여러분이 정합니다(층 2 이상 8 이하, 헤드 2 이상 8 이하, 헤드 차원 4 이상 32 이하의 짝수, 자료형 바이트는 1·2·4 중 하나). mul 은 곱하면서 계수기를 1 올리고, dot 은 그 mul 로만 곱합니다.

계수기는 모듈 안의 정수 하나면 됩니다. 함수 안에서 바꾸려면 global 이 필요합니다. dotsum(a * b for a, b in zip(u, v)) 처럼 직접 곱하면 아무것도 세지 않게 되니 반드시 mul 을 거치세요. 덧셈은 세지 않습니다 — 행렬 계산의 값은 곱셈 쪽에서 나옵니다. 시간을 재는 코드는 넣지 마세요.

토큰 하나를 투영한다

project(x, W) 를 더하세요. W 는 줄마다 길이가 len(x) 인 행렬이고, 줄 하나가 출력 한 칸을 만듭니다. 돌려주는 목록의 길이는 len(W) 이고 곱셈은 len(W) 곱하기 len(x) 번 일어납니다.

앞에서 만든 dot 을 줄마다 한 번씩 부르면 끝입니다. 한 줄로 적을 수 있습니다. 줄과 칸을 바꿔 쓰면 정사각 행렬에서는 값만 틀리고 직사각에서는 길이까지 어긋나니, 돌려주는 목록의 길이가 len(W) 인지 먼저 확인하세요. q 도 k 도 v 도 이 함수 하나로 만듭니다.

질의 하나가 캐시 전체를 본다

attend(q, K, V) 를 더하세요. q 와 각 k 의 내적을 √len(q) 로 나눠 점수를 내고, 최댓값을 뺀 뒤 지수를 취해 소프트맥스를 씌우고, 그 가중치로 V 의 가중 평균을 냅니다. 곱셈은 점수 쪽 len(K) 곱하기 len(q) 번, 가중합 쪽 len(V) 곱하기 len(V[0]) 번입니다.

나누기와 지수는 곱셈이 아니므로 mul 로 감싸지 마세요 — 감싸면 횟수가 어긋납니다. 가중합의 곱셈만 mul 을 거칩니다. V 한 줄의 길이가 q 의 길이와 다를 수 있으니 출력 목록은 len(V[0]) 으로 잡으세요. K 가 한 줄뿐이면 가중치가 1 하나라 출력이 V[0] 과 같아야 합니다.

캐시 없이 — 앞 전체를 다시 계산한다

step_nocache(xs, Wq, Wk, Wv) 를 더하세요. xs모든 토큰에 대해 q·k·v 를 만들고, 마지막 질의로 전체 K·V 를 봅니다. 쓰이는 질의는 하나뿐인데도 전부 만드는 것이 캐시 없는 상태의 모습입니다.

세 번의 목록 내포와 attend 한 번이면 됩니다. 마지막 질의만 쓴다고 해서 마지막 토큰만 투영하면 안 됩니다 — 키와 값은 앞 토큰 전부에 대해 있어야 하고, 이 단계의 요점은 그것을 매번 다시 만든다는 사실입니다. 곱셈 횟수는 문맥 길이에 비례해 늘어납니다.

캐시로 — 한 줄만 덧붙인다

new_cache()·store(cache, k, v, digits=None)·step_cached(x, Wq, Wk, Wv, cache, digits=None) 를 더하세요. new_cache(){"K": [], "V": []} 를 돌려주고, store 는 K 와 V 에 각각 한 줄 덧붙이며(digits 가 주어지면 그 소수 자릿수로 반올림해서), step_cached 는 새 토큰 하나만 투영해 캐시에 넣은 그 질의로 캐시 전체를 봅니다.

덮어쓰지 말고 append 하세요. 앞 줄들은 새 토큰이 붙어도 값이 바뀌지 않습니다 — 인과 마스크 때문에 각 자리가 자기 앞만 보기 때문이고, 그것이 캐시가 성립하는 이유입니다. 넣기 전에 어텐션하면 새 토큰이 자기 자신을 못 보게 되어 캐시 없는 판과 답이 달라집니다. 곱셈 횟수는 투영 쪽이 문맥 길이와 무관하게 고정됩니다.

두 방식의 출력이 같은 값인가

ATOL = 1e-9·RTOL = 1e-6·close(got, want)·compare(xs, Wq, Wk, Wv, digits=None) 를 더하세요. closeabs(got - want) <= ATOL + RTOL * abs(want) 이고 왜 이 폭인지 주석으로 적습니다. compare 는 길이 1 부터 한 칸씩 늘려 가며 두 방식을 돌리고 {"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]} 를 돌려줍니다.

등호로 견주지 마세요. 두 방식이 우연히 같은 순서로 더하면 비트까지 같게 나올 수도 있지만, 더하는 순서가 조금만 달라져도 마지막 자리가 흔들립니다 — 그래서 옳은 시험은 허용 오차입니다. digitsstep_cached 로 그대로 넘깁니다. 반올림해 넣으면 all_close 가 거짓이 되는데, 그게 이 단계에서 보려는 것입니다. max_gap 은 모든 길이·모든 성분을 통틀어 가장 큰 차이입니다.

길이마다 곱셈을 센다

mul_table(xs, Wq, Wk, Wv) 를 더하세요. 문맥 길이 1 부터 len(xs) 까지 두 방식의 곱셈 횟수를 세어 (문맥길이, 캐시없음, 캐시있음) 짝의 목록을 돌려줍니다. 재기 직전마다 계수기를 되돌리고, 캐시 쪽은 한 벌의 캐시를 계속 이어 씁니다.

reset_muls() 를 두 번 부릅니다 — 캐시 없는 판을 재기 전에 한 번, 캐시 판을 재기 전에 한 번. 되돌리지 않으면 뒤의 수에 앞의 수가 섞여 들어옵니다. 캐시를 줄마다 새로 만들면 캐시 쪽 수가 캐시 없는 쪽처럼 자라 버립니다 — 그건 캐시가 아닙니다. 표를 보면 캐시 없는 쪽은 길이에 비례해 늘고, 캐시 쪽은 어텐션 몫만 늘어납니다.

아끼는 것과 치르는 것을 함께 적는다

d = CONFIG["head_dim"] 로 두고 토큰 12개, d 곱하기 d 짜리 가중치 셋으로 mul_tablecompare 를 돌리세요. 값은 ((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0 으로 만들고 salt 는 토큰 벡터 1, Wq 2, Wk 4, Wv 6 입니다. compare 는 반올림 없이 한 번, digits=2 로 한 번 돌립니다. 그리고 /root/work/tf-kv/kv_report.jsonlayers·heads·head_dim·dtype_bytes·tokens·table·total_nocache·total_cached·saved_muls·cache_elems·cache_bytes·round_digits·max_gap_exact·all_close_exact·max_gap_rounded·all_close_rounded 를, /root/work/tf-kv/kv_report.md## 무엇을 쟀나 ## 캐시 없이 하면 무엇을 다시 계산하나 ## 캐시가 먹는 메모리 ## 두 방식의 출력이 같은가 네 절로 쓰세요.

숫자는 손으로 적지 말고 여러분의 코드를 실제로 돌려 얻은 값으로 채우세요. total_nocache·total_cached 는 표의 각 칸을 더한 것이고 saved_muls 는 그 차이입니다. cache_elems2 × layers × heads × 12 × head_dim, cache_bytes 는 거기에 dtype_bytes 를 곱한 값입니다. all_close_exact 는 참, all_close_rounded 는 거짓이어야 합니다 — 캐시 자체는 근사가 아니지만 캐시에 덜 정확하게 적는 것은 근사입니다. 값 만드는 식에서 7 로 나누는 것은 값이 소수 둘째 자리에서 딱 떨어지지 않게 하려는 것입니다. 딱 떨어지는 값만 쓰면 반올림해도 값이 그대로라 이 실험 자체가 성립하지 않습니다. 시간은 재지 마세요.