트랜스포머 — 어텐션을 손으로 계산한다 · KV 캐시가 줄이는 계산량 · 실습
KV 캐시가 아끼는 곱셈을 센다
목표
자기회귀 생성에서 KV 캐시가 무엇을 아끼는지 곱셈 횟수를 직접 세어 확인한다. 곱셈을 세는 곱셈 함수를 만들고, 캐시 없이 도는 판과 캐시로 도는 판을 각각 만들어 두 판의 출력이 같은 값인지 허용 오차로 확인한 뒤, 문맥 길이마다 곱셈 횟수를 세어 표로 만든다. 마지막에 캐시를 낮은 정밀도로 들고 있으면 그 같음이 깨지는 것을 재고, 자기가 정한 설정으로 캐시가 먹는 원소 수와 바이트를 계산해 기록으로 남긴다.
왜 중요한가
어텐션 한 번을 손으로 계산하는 것과, 그 어텐션을 생성 반복 안에서 부르는 것은 다른 이야기다. 길이 100짜리 답을 만들려면 모델을 100번 부르고, 캐시가 없으면 부를 때마다 앞 문맥 전체의 k·v 를 처음부터 다시 만든다. 첫 토큰의 k·v 를 백 번 만들고 아흔아홉 번 버린다는 뜻이다.
이 낭비는 코드를 읽어서는 보이지 않는다. 어텐션 함수 자체는 잘못된 곳이 하나도 없기 때문이다. 보이게 하려면 세어야 한다. 그리고 시간을 재서는 안 된다 — 시간은 기계와 부하에 따라 달라지지만 곱셈 횟수는 같은 입력에서 늘 같고, 길이에 대해 어떻게 늘어나는지를 그대로 보여 준다.
이 실습은 실제 모델을 부르지 않는다. 이 파드의 시스템 파이썬에는 numpy 가 없고(/opt/onnx-lab/bin/python 안에만 있다) torch 도 transformers 도 없다. 표준 라이브러리만으로 같은 구조를 만들고, 여기서 잰 숫자만 쓴다. 그래서 "실제 모델은 몇 GB 를 먹는다" 나 "몇 배 빨라진다" 같은 말은 여기서 하지 않는다.
채점기는 여러분이 적어 둔 설명을 믿지 않는다. 여러분의 모듈을 실제로 불러 매번 다른 크기의 입력으로 함수를 두드려 보고, 값과 곱셈 횟수를 따로 계산해 대조한다. 크기가 실행마다 바뀌므로 값을 외워 넣을 수 없다.
단계
1. /root/work/tf-kv/kv.py 에 CONFIG 와 reset_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.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](https://arxiv.org/abs/1706.03762) · [Hugging Face — Cache strategies](https://huggingface.co/docs/transformers/en/kv_cache) · [Hugging Face — Text generation](https://huggingface.co/docs/transformers/en/llm_tutorial) · [Python — math](https://docs.python.org/3/library/math.html)
- 흔한 실수:
mul을 거치지 않고 직접 곱해 셈이 0 이 되기,project에서 줄과 칸을 바꿔 쓰기, 점수를√d로 나누지 않기, 캐시에 덮어쓰기, 캐시에 넣기 전에 어텐션하기,mul_table에서 계수기를 안 되돌리기, 캐시를 줄마다 새로 만들기.
단계 8개
- 곱셈을 세는 곱셈
- 토큰 하나를 투영한다
- 질의 하나가 캐시 전체를 본다
- 캐시 없이 — 앞 전체를 다시 계산한다
- 캐시로 — 한 줄만 덧붙인다
- 두 방식의 출력이 같은 값인가
- 길이마다 곱셈을 센다
- 아끼는 것과 치르는 것을 함께 적는다