트랜스포머 — 어텐션을 손으로 계산한다 · 문맥 길이와 비용 · 실습
문맥 길이의 비용을 직접 센다
목표
곱셈을 세는 계수기를 만들어 문맥 길이 n 이 늘 때 무엇이 n 제곱으로 늘고 무엇이 n 에 비례해 느는지 직접 세어 확인한다. 두 몫을 나란히 놓은 표를 만들고, n 제곱 몫이 선형 몫을 따라잡는 교차점을 자기가 센 숫자로 찾는다. 인과 마스크가 남기는 점수가 n(n+1)/2 개뿐이라는 것, 점수 행렬을 통째로 들면 원소가 n 제곱 개라는 것, 한 토큰을 붙일 때 새로 생기는 점수는 줄 하나뿐이라는 것까지 정수로 센다.
왜 중요한가
"문맥 길이에 대해 제곱" 이라는 말은 반쪽만 맞다. 한 층 안에는 성질이 다른 두 계산이 섞여 있다. 어텐션의 점수와 섞기는 자리마다 문맥 전체를 훑어서 제곱으로 자라지만, 질의·열쇠·값 투영과 출력 투영과 피드포워드는 자리 하나하나에 똑같이 한 번씩 도는 일이라 자리 수에만 비례한다. 어느 쪽이 지배하는지는 길이와 모델 폭이 함께 정한다. 그래서 문맥 한도를 늘렸을 때의 비용을 하나의 배율로 어림하면 반드시 틀린다.
이 실습은 시간을 재지 않는다. 같은 코드도 기계와 부하에 따라 다르고 컨테이너 안에서는 더 흔들려서, 재어 봐야 견줄 수가 없다. 대신 곱셈 횟수를 센다. 횟수는 어디서 돌려도 같은 정수라 남에게 근거로 보여 줄 수 있다. 여기 나오는 숫자는 전부 여러분의 계수기가 센 값이고, 실제 모델의 초나 GB 는 재지 않으므로 쓰지 않는다.
모델 모양은 Attention Is All You Need 의 base 설정을 가정으로 쓴다 — D_MODEL = 512, D_FF = 2048, N_HEADS = 8. 이것은 이 실습이 정한 가정이지 여러분이 쓰는 모델의 값이 아니다.
채점기는 여러분이 적어 둔 설명을 믿지 않는다. 여러분의 모듈을 실제로 불러 매번 다른 크기로 함수를 두드려 보고, 계수기가 실제로 얼마나 올라갔는지까지 대조한다. 크기는 실행마다 바뀌므로 값을 외워 넣을 수 없다.
단계
1. /root/work/tf-cost/cost.py 에 상수 D_MODEL = 512·D_FF = 2048·N_HEADS = 8 과 계수기 MulCount, 그리고 그것을 쓰는 dot(a, b, ctr) 를 만드세요.
2. attn_scores(Q, K, ctr)·attn_mix(A, V, ctr)·quad_mults(n, d_model) 을 더해 n 제곱으로 자라는 몫을 세게 하세요.
3. matvec(M, x, ctr)·per_position_mults(d_model, d_ff, ctr)·linear_mults(n, d_model, d_ff) 를 더해 자리 수에 비례하는 몫을 세게 하세요.
4. cost_table(ns, d_model, d_ff) 를 만들어 길이마다 (n, n제곱 몫, 선형 몫, 합) 을 돌려주게 하세요.
5. crossover_n(d_model, d_ff) 를 만들어 n 제곱 몫이 선형 몫을 처음으로 따라잡는 길이를 찾게 하세요.
6. causal_pairs(n) 과 wasted_pairs(n) 을 만들어 인과 마스크가 남기는 점수와 버리는 칸을 세게 하세요.
7. score_bytes(n, n_heads, itemsize)·max_context_for_bytes(budget_bytes, n_heads, itemsize)·append_pairs(n)·generate_pairs(n, g) 를 만드세요.
8. 위 함수들을 실제로 돌린 결과를 /root/work/tf-cost/cost_report.json 과 /root/work/tf-cost/cost_report.md 에 기록하세요.
참고
- 실행 계약: 채점기는
/root/work/tf-cost/cost.py를 파이썬 모듈로 불러D_MODEL·D_FF·N_HEADS·MulCount·dot·attn_scores·attn_mix·quad_mults·matvec·per_position_mults·linear_mults·cost_table·crossover_n·causal_pairs·wasted_pairs·score_bytes·max_context_for_bytes·append_pairs·generate_pairs를 직접 씁니다. 스크립트로 실행하지 않으므로if __name__ == "__main__"은 없어도 됩니다. MulCount는mults(지금까지 센 곱셈 횟수)와calls(장부에 적은 횟수) 두 값을 들고,add(k)로 k 를 더합니다. 처음에는 둘 다 0 입니다.dot(a, b, ctr)는 내적 값을 돌려주고 계수기를 정확히 길이만큼 올립니다. 곱셈 한 번을 1 로 세는 것이지 호출 한 번을 1 로 세는 것이 아닙니다.attn_scores(Q, K, ctr)는 Q 의 줄 수 x K 의 줄 수 크기의 행렬을 돌려줍니다. Q 와 K 가 각각 n x d 면 계수기는n * n * d만큼 올라갑니다.attn_mix(A, V, ctr)는 A 가 n x n, V 가 n x d 일 때 n x d 를 돌려주고 계수기는 또n * n * d만큼 올라갑니다. 출력이 작다고 계산이 작은 것이 아닙니다.quad_mults(n, d_model)은 두 몫을 합쳐2 * n * n * d_model입니다. 헤드 수는 들어가지 않습니다 — 헤드마다 폭이d_model / h로 줄고 그런 헤드가 h 개라 합이 같기 때문입니다.per_position_mults(d_model, d_ff, ctr)는 선형 사상 여섯 개를 각각 실제로 돌려 곱셈을 셉니다. 질의·열쇠·값 투영 셋(d_model x d_model), 출력 투영 하나(d_model x d_model), 피드포워드 1층(d_ff x d_model)과 2층(d_model x d_ff) 입니다. 돌려주는 값은 이 함수가 올린 만큼이고, 채점기는 계수기가 최소 여섯 번 이상 불렸는지도 봅니다. 식 하나를 한 번에 더하고 끝내면 떨어집니다.linear_mults(n, d_model, d_ff)는 자리 하나의 값에 자리 수를 곱한 정수입니다. 계수기를 받지 않습니다.cost_table(ns, d_model, d_ff)가 돌려주는 줄은(n, quad, linear, quad + linear)네 칸이고 순서는ns와 같습니다.crossover_n(d_model, d_ff)는quad_mults(n, d_model) >= linear_mults(n, d_model, d_ff)가 처음으로 참이 되는 n 입니다. 등호를 포함합니다. 1 부터 올려 가며 찾아도 되고 식을 풀어도 됩니다.causal_pairs(n)은 i 번째 줄이i + 1개를 쓴다는 사실에서 나옵니다.causal_pairs(0)은 0,causal_pairs(1)은 1 입니다.wasted_pairs(n)은n * n에서 실제로 쓰는 몫을 뺀 값입니다.score_bytes(n, n_heads, itemsize)는n * n * n_heads * itemsize입니다. 헤드 수를 빠뜨리지 마세요.max_context_for_bytes(budget_bytes, n_heads, itemsize)는score_bytes(n, ...) <= budget_bytes를 만족하는 가장 큰 n 입니다. 실수 제곱근을 반올림하면 한 칸 넘치는 일이 생기므로math.isqrt를 쓰세요.append_pairs(n)은n + 1,generate_pairs(n, g)는 줄 길이n+1부터n+g까지의 합입니다.generate_pairs(0, N)이causal_pairs(N)과 같아야 합니다.- 8단계 보고서는
D_MODEL·D_FF·N_HEADS와ITEMSIZE = 2(두 바이트 자료형을 가정), 예산1073741824(1 GiB), 표의 길이[128, 256, 512, 1024, 2048, 4096], 인과·메모리 계산의 기준 길이2048, 생성 토큰 수256을 씁니다. quad_ratio·linear_ratio는 표의 마지막 두 줄(2048 과 4096) 사이의 배율입니다. 나눗셈이므로 실수이고, 채점기는abs(a - b) <= atol + rtol * abs(b)로 견줍니다.- 이 파드에는 인터넷이 없습니다.
pip install은 되지 않고 시스템 파이썬에는 numpy·torch·transformers 가 없습니다. numpy 는/opt/onnx-lab/bin/python안에만 있습니다. 표준 라이브러리만으로 충분합니다. - 공식 문서: [Attention Is All You Need](https://arxiv.org/abs/1706.03762) · [PyTorch — scaled_dot_product_attention](https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html) · [Hugging Face — Cache strategies](https://huggingface.co/docs/transformers/en/kv_cache)
- 흔한 실수: 계수기를 호출 횟수로 세기,
attn_mix가 출력 크기만큼만 센다고 생각하기, 선형 몫에서 출력 투영을 빠뜨리기, 교차점을 등호 없이 찾아 한 칸 밀리기,causal_pairs에서 대각선을 빼기,score_bytes에 헤드 수를 안 곱하기.
무엇을 재지 않는가
시간을 재지 않습니다. 실제 모델의 GB 나 초도 쓰지 않습니다. 재어 보지 않은 숫자를 기록에 적으면 그 기록은 근거가 없습니다.
단계 8개
- 곱셈을 세는 자를 만든다
- n 제곱으로 자라는 몫
- 자리 수에만 비례하는 몫
- 두 몫을 나란히 놓는다
- 교차점을 찾는다
- 마스크가 버리는 절반을 센다
- 메모리와 한 줄씩 늘어나는 점수
- 센 숫자를 기록으로 남긴다