KV キャッシュが節約する乗算を数える
한국어 원문으로 표시합니다.
목표
자기회귀 생성에서 KV 캐시가 무엇을 아끼는지 곱셈 횟수를 직접 세어 확인한다. 곱셈을 세는 곱셈 함수를 만들고, 캐시 없이 도는 판과 캐시로 도는 판을 각각 만들어 두 판의 출력이 같은 값인지 허용 오차로 확인한 뒤, 문맥 길이마다 곱셈 횟수를 세어 표로 만든다. 마지막에 캐시를 낮은 정밀도로 들고 있으면 그 같음이 깨지는 것을 재고, 자기가 정한 설정으로 캐시가 먹는 원소 수와 바이트를 계산해 기록으로 남긴다.
왜 중요한가
어텐션 한 번을 손으로 계산하는 것과, 그 어텐션을 생성 반복 안에서 부르는 것은 다른 이야기다. 길이 100짜리 답을 만들려면 모델을 100번 부르고, 캐시가 없으면 부를 때마다 앞 문맥 전체의 k·v 를 처음부터 다시 만든다. 첫 토큰의 k·v 를 백 번 만들고 아흔아홉 번 버린다는 뜻이다.
이 낭비는 코드를 읽어서는 보이지 않는다. 어텐션 함수 자체는 잘못된 곳이 하나도 없기 때문이다. 보이게 하려면 세어야 한다. 그리고 시간을 재서는 안 된다 — 시간은 기계와 부하에 따라 달라지지만 곱셈 횟수는 같은 입력에서 늘 같고, 길이에 대해 어떻게 늘어나는지를 그대로 보여 준다.
이 실습은 실제 모델을 부르지 않는다. 이 파드의 시스템 파이썬에는 numpy 가 없고(/opt/onnx-lab/bin/python 안에만 있다) torch 도 transformers 도 없다. 표준 라이브러리만으로 같은 구조를 만들고, 여기서 잰 숫자만 쓴다. 그래서 "실제 모델은 몇 GB 를 먹는다" 나 "몇 배 빨라진다" 같은 말은 여기서 하지 않는다.
채점기는 여러분이 적어 둔 설명을 믿지 않는다. 여러분의 모듈을 실제로 불러 매번 다른 크기의 입력으로 함수를 두드려 보고, 값과 곱셈 횟수를 따로 계산해 대조한다. 크기가 실행마다 바뀌므로 값을 외워 넣을 수 없다.
단계
- /root/work/tf-kv/kv.py 에
CONFIG와reset_muls()·muls()·mul(a, b)·dot(u, v)를 만드세요.mul은 곱하면서 센 횟수를 하나 올리고,dot은 그mul로만 곱합니다. project(x, W)를 더해 토큰 하나를 가중치 한 벌에 통과시키게 하세요.W는 줄마다 길이가len(x)인 행렬이고 돌려주는 목록의 길이는len(W)입니다.attend(q, K, V)를 더해 질의 하나가 쌓인 K·V 전체를 보게 하세요. 점수를√len(q)로 나누고 소프트맥스를 씌운 뒤 V 의 가중 평균을 냅니다.step_nocache(xs, Wq, Wk, Wv)를 더해 캐시 없이 출력 하나를 내게 하세요. 앞의 모든 토큰에 대해 q·k·v 를 다시 만들고 마지막 질의로 어텐션합니다.new_cache()·store(cache, k, v, digits=None)·step_cached(x, Wq, Wk, Wv, cache, digits=None)를 더해 캐시로 같은 출력을 내게 하세요. 새 토큰 하나만 계산하고 K·V 에 한 줄 덧붙입니다.ATOL·RTOL·close(got, want)·compare(xs, Wq, Wk, Wv, digits=None)를 더해 두 방식의 출력이 같은 값인지 견주게 하세요. 등호로 견주지 않습니다.mul_table(xs, Wq, Wk, Wv)를 더해 문맥 길이 1 부터len(xs)까지 두 방식의 곱셈 횟수를 세게 하세요. 돌려주는 값은(문맥길이, 캐시없음, 캐시있음)짝의 목록입니다.- 토큰 12개로 표를 만들고 캐시가 먹는 메모리를 계산해 /root/work/tf-kv/kv_report.json 과 /root/work/tf-kv/kv_report.md 에 기록하세요.
참고
- 실행 계약: 채점기는
/root/work/tf-kv/kv.py를 파이썬 모듈로 불러CONFIG·reset_muls·muls·mul·dot·project·attend·step_nocache·new_cache·store·step_cached·ATOL·RTOL·close·compare·mul_table을 직접 씁니다. 스크립트로 실행하지 않으므로if __name__ == "__main__"은 없어도 됩니다. CONFIG는{"layers": ..., "heads": ..., "head_dim": ..., "dtype_bytes": ...}네 열쇠를 가진 딕셔너리입니다. 값은 여러분이 정합니다. 범위는 층 2 이상 8 이하, 헤드 2 이상 8 이하, 헤드 차원 4 이상 32 이하의 짝수, 자료형 바이트는 1·2·4 중 하나입니다. 실제 모델의 수를 흉내 낼 필요가 없습니다 — 8단계의 메모리 계산은 여러분이 정한 이 값으로만 합니다.mul(a, b)는a * b를 돌려주면서 모듈 안의 계수기를 1 올립니다.reset_muls()는 계수기를 0 으로,muls()는 지금 값을 돌려줍니다. 곱셈이 일어나는 자리를 전부mul로 통과시켜야 셈이 맞습니다.dot(u, v)는 내적입니다. 곱셈은len(u)번 일어납니다. 덧셈·나눗셈·지수는 세지 않습니다.project(x, W)의 곱셈은len(W)곱하기len(x)번입니다.W의 줄이 출력 한 칸을 만듭니다 — 줄과 칸을 바꿔 쓰면 값도 횟수도 어긋납니다.attend(q, K, V)의 곱셈은 점수 쪽len(K)곱하기len(q)번, 가중합 쪽len(V)곱하기len(V[0])번입니다. 나누기와 지수는 곱셈이 아니므로 세지 않습니다. 소프트맥스는 최댓값을 뺀 뒤 지수를 취하세요.step_nocache(xs, Wq, Wk, Wv)는xs의 모든 토큰에 대해 q·k·v 를 만들고, 마지막 질의로 전체 K·V 를 봅니다. 쓰이는 질의가 하나뿐인데도 전부 만드는 것이 캐시 없는 상태의 모습입니다.new_cache()는{"K": [], "V": []}를 돌려줍니다.store(cache, k, v, digits=None)는 K 와 V 에 각각 한 줄 덧붙입니다 — 덮어쓰지 않습니다.digits가 주어지면 각 성분을 그 소수 자릿수로 반올림해 넣습니다.step_cached(x, Wq, Wk, Wv, cache, digits=None)는 새 토큰 하나만 투영하고,store로 캐시에 넣은 뒤에, 그 질의로 캐시 전체를 봅니다. 넣기 전에 어텐션하면 새 토큰이 자기 자신을 보지 못합니다.ATOL = 1e-9,RTOL = 1e-6으로 두고close(got, want)는abs(got - want) <= ATOL + RTOL * abs(want)입니다. 왜 이 폭인지 주석으로 적으세요. 파이썬 float 는 IEEE 754 배정밀도라 유효숫자가 약 15자리이고, 이 규모의 내적·소프트맥스를 순서만 바꿔 더해도 상대 오차는 1e-12 수준에 머뭅니다.compare(xs, Wq, Wk, Wv, digits=None)는 길이 1 부터len(xs)까지 한 칸씩 늘려 가며 두 방식을 돌리고{"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]}를 돌려줍니다.max_gap은 성분끼리 차이의 절댓값 중 가장 큰 값이고,all_close는 모든 성분이close를 통과했는지입니다.mul_table(xs, Wq, Wk, Wv)는 길이마다 재기 직전에 계수기를 되돌립니다. 캐시 쪽은 한 벌의 캐시를 계속 이어 써야 합니다 — 줄마다 새로 만들면 그건 캐시가 아닙니다.- 8단계는
d = CONFIG["head_dim"]로 두고 토큰 12개,d곱하기d짜리 가중치 셋으로 잽니다.round_digits는 2 를 씁니다. - 8단계의 토큰 벡터와 가중치는
((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0로 만듭니다.i는 줄 번호,j는 칸 번호이고salt는 토큰 벡터 1, Wq 2, Wk 4, Wv 6 입니다. 곱셈 횟수와 메모리는 값이 아니라 크기에서만 나오므로 어떤 값을 써도 표는 같지만, 반올림 실험은 값에 달려 있습니다 — 7 로 나누는 것은 값이 소수 둘째 자리에서 딱 떨어지지 않게 하려는 것입니다. 딱 떨어지는 값만 쓰면 소수 둘째 자리로 반올림해도 값이 그대로라 반올림의 영향이 보이지 않습니다. - 메모리는
원소 수 = 2 × layers × heads × 길이 × head_dim,바이트 = 원소 수 × dtype_bytes입니다. 길이는 12 입니다. - 이 파드에는 인터넷이 없습니다.
pip install은 되지 않고 torch·transformers 도 없습니다. numpy 는/opt/onnx-lab/bin/python안에만 있어 시스템 파이썬에서는import numpy가 되지 않습니다.math만으로 충분합니다. - 시간을 재지 마세요.
time으로 잰 수는 기계와 부하에 따라 달라져 판정에 쓸 수 없습니다. 이 실습이 재는 것은 곱셈 횟수입니다. - 공식 문서: Attention Is All You Need · Hugging Face — Cache strategies · Hugging Face — Text generation · Python — math
- 흔한 실수:
mul을 거치지 않고 직접 곱해 셈이 0 이 되기,project에서 줄과 칸을 바꿔 쓰기, 점수를√d로 나누지 않기, 캐시에 덮어쓰기, 캐시에 넣기 전에 어텐션하기,mul_table에서 계수기를 안 되돌리기, 캐시를 줄마다 새로 만들기.
곱셈을 세는 곱셈
/root/work/tf-kv/kv.py 에 CONFIG 와 reset_muls()·muls()·mul(a, b)·dot(u, v) 를 만드세요. CONFIG 는 layers·heads·head_dim·dtype_bytes 네 열쇠를 가진 딕셔너리이고 값은 여러분이 정합니다(층 2 이상 8 이하, 헤드 2 이상 8 이하, 헤드 차원 4 이상 32 이하의 짝수, 자료형 바이트는 1·2·4 중 하나). mul 은 곱하면서 계수기를 1 올리고, dot 은 그 mul 로만 곱합니다.
계수기는 모듈 안의 정수 하나면 됩니다. 함수 안에서 바꾸려면 global 이 필요합니다. dot 이 sum(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) 를 더하세요. close 는 abs(got - want) <= ATOL + RTOL * abs(want) 이고 왜 이 폭인지 주석으로 적습니다. compare 는 길이 1 부터 한 칸씩 늘려 가며 두 방식을 돌리고 {"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]} 를 돌려줍니다.
등호로 견주지 마세요. 두 방식이 우연히 같은 순서로 더하면 비트까지 같게 나올 수도 있지만, 더하는 순서가 조금만 달라져도 마지막 자리가 흔들립니다 — 그래서 옳은 시험은 허용 오차입니다. digits 는 step_cached 로 그대로 넘깁니다. 반올림해 넣으면 all_close 가 거짓이 되는데, 그게 이 단계에서 보려는 것입니다. max_gap 은 모든 길이·모든 성분을 통틀어 가장 큰 차이입니다.
길이마다 곱셈을 센다
mul_table(xs, Wq, Wk, Wv) 를 더하세요. 문맥 길이 1 부터 len(xs) 까지 두 방식의 곱셈 횟수를 세어 (문맥길이, 캐시없음, 캐시있음) 짝의 목록을 돌려줍니다. 재기 직전마다 계수기를 되돌리고, 캐시 쪽은 한 벌의 캐시를 계속 이어 씁니다.
reset_muls() 를 두 번 부릅니다 — 캐시 없는 판을 재기 전에 한 번, 캐시 판을 재기 전에 한 번. 되돌리지 않으면 뒤의 수에 앞의 수가 섞여 들어옵니다. 캐시를 줄마다 새로 만들면 캐시 쪽 수가 캐시 없는 쪽처럼 자라 버립니다 — 그건 캐시가 아닙니다. 표를 보면 캐시 없는 쪽은 길이에 비례해 늘고, 캐시 쪽은 어텐션 몫만 늘어납니다.
아끼는 것과 치르는 것을 함께 적는다
d = CONFIG["head_dim"] 로 두고 토큰 12개, d 곱하기 d 짜리 가중치 셋으로 mul_table 과 compare 를 돌리세요. 값은 ((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.json 에 layers·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_elems 는 2 × layers × heads × 12 × head_dim, cache_bytes 는 거기에 dtype_bytes 를 곱한 값입니다. all_close_exact 는 참, all_close_rounded 는 거짓이어야 합니다 — 캐시 자체는 근사가 아니지만 캐시에 덜 정확하게 적는 것은 근사입니다. 값 만드는 식에서 7 로 나누는 것은 값이 소수 둘째 자리에서 딱 떨어지지 않게 하려는 것입니다. 딱 떨어지는 값만 쓰면 반올림해도 값이 그대로라 이 실험 자체가 성립하지 않습니다. 시간은 재지 마세요.