LabHub
배우기 러닝패스 코스

Transformers — Compute Attention By Hand

Count the Cost of Context Length Yourself

LabHub 에서 이어서 보기

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

목표

곱셈을 세는 계수기를 만들어 문맥 길이 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 에 기록하세요.

참고

무엇을 재지 않는가

시간을 재지 않습니다. 실제 모델의 GB 나 초도 쓰지 않습니다. 재어 보지 않은 숫자를 기록에 적으면 그 기록은 근거가 없습니다.

곱셈을 세는 자를 만든다

/root/work/tf-cost/cost.py 에 상수 D_MODEL = 512·D_FF = 2048·N_HEADS = 8 과 계수기 MulCount(mults·calls 를 들고 add(k) 로 올린다), 그리고 dot(a, b, ctr) 를 만드세요. dot 은 내적 값을 돌려주고 계수기를 정확히 벡터 길이만큼 올립니다.

시간을 재려 하지 마세요 — 같은 코드도 기계와 부하에 따라 달라져 견줄 수가 없습니다. 길이 d 인 내적은 곱셈이 정확히 d 번이니 ctr.add(len(a)) 한 줄이면 됩니다. 호출 한 번을 1 로 세면 뒤의 모든 숫자가 무너집니다. calls 는 장부에 몇 번 나누어 적었는지를 세는 값으로, 3단계에서 쓰입니다.

n 제곱으로 자라는 몫

attn_scores(Q, K, ctr)·attn_mix(A, V, ctr)·quad_mults(n, d_model) 을 더하세요. attn_scores 는 n x n 점수 행렬을, attn_mix 는 A 로 V 를 섞은 n x d 를 돌려주고 둘 다 계수기를 n * n * d 만큼 올립니다. quad_mults 는 둘을 합친 2 * n * n * d_model 입니다.

attn_scores 는 질의마다 열쇠 전부와 dot 하면 끝입니다. attn_mix 가 헷갈리는 자리입니다 — 나오는 것은 n x d 로 작지만 나오는 자리마다 문맥 전체 n 개를 훑어야 하므로 곱셈은 점수 계산과 똑같이 n * n * d 번입니다. V 의 세로줄을 뽑아 dot 에 넘기면 계수기가 알아서 맞습니다. quad_mults 에 헤드 수는 들어가지 않습니다.

자리 수에만 비례하는 몫

matvec(M, x, ctr)·per_position_mults(d_model, d_ff, ctr)·linear_mults(n, d_model, d_ff) 를 더하세요. per_position_mults 는 선형 사상 여섯 개(투영 넷, 피드포워드 둘)를 각각 실제로 돌려 곱셈을 세고 올린 만큼을 돌려줍니다. linear_mults 는 그 값에 자리 수를 곱한 정수입니다.

행렬의 값은 아무래도 좋습니다 — 세는 것이 목적이므로 모양만 맞춘 행렬을 만들어 돌리면 됩니다. 여섯 개는 d_model x d_model 넷과 d_ff x d_model 하나, d_model x d_ff 하나입니다. 피드포워드 2층에 넘기는 벡터는 1층이 돌려준 길이 d_ff 짜리입니다. 채점기는 계수기가 최소 여섯 번 이상 불렸는지도 보므로, 식 하나를 한 번에 더하고 끝내면 떨어집니다. linear_mults 는 계수기를 받지 않습니다.

두 몫을 나란히 놓는다

cost_table(ns, d_model, d_ff) 를 만드세요. ns 의 길이마다 (n, quad, linear, quad + linear) 네 칸을 돌려주고 순서는 ns 와 같습니다.

앞에서 만든 quad_multslinear_mults 를 그대로 쓰면 다섯 줄입니다. 길이를 두 배씩 늘리며 읽어 보세요 — 앞쪽 칸은 네 배씩, 뒤쪽 칸은 두 배씩 갑니다. 마지막 칸은 반드시 앞 두 칸의 합이어야 합니다. 한쪽만 적어 두면 뒤에서 교차점을 찾을 때 어긋납니다.

교차점을 찾는다

crossover_n(d_model, d_ff) 를 만드세요. quad_mults(n, d_model) >= linear_mults(n, d_model, d_ff)처음으로 참이 되는 n 입니다. 등호를 포함합니다.

1 부터 올려 가며 찾아도 되고 손으로 풀어도 됩니다. 양변을 n2 * d_model 로 나누면 조건이 아주 짧아집니다. 등호를 빠뜨리고 부등호만 쓰면 답이 정확히 한 칸 밀립니다. 모델 폭을 바꿔 가며 불러 보세요 — 폭이 넓을수록 교차점이 뒤로 밀린다는 것이 숫자로 보입니다.

마스크가 버리는 절반을 센다

causal_pairs(n)wasted_pairs(n) 을 만드세요. causal_pairs 는 인과 마스크에서 실제로 쓰이는 점수의 수이고, wasted_pairsn * n 에서 그 몫을 뺀 값입니다.

i 번째 질의는 자기 자신까지만 봅니다. 줄 길이가 1, 2, 3 으로 늘어 n 에서 끝나므로 세어 보면 곧 n * (n + 1) / 2 입니다. 대각선을 빼면 안 됩니다 — 자기 자신은 보는 것이 맞습니다. causal_pairs(0) 은 0, causal_pairs(1) 은 1 입니다. 버려지는 비율이 길이가 길어질수록 어디로 다가가는지 몇 개 찍어 보세요.

메모리와 한 줄씩 늘어나는 점수

score_bytes(n, n_heads, itemsize)·max_context_for_bytes(budget_bytes, n_heads, itemsize)·append_pairs(n)·generate_pairs(n, g) 를 만드세요. 앞의 둘은 점수 행렬을 통째로 들 때의 바이트와 예산 안에 드는 가장 긴 문맥이고, 뒤의 둘은 한 토큰을 붙일 때 새로 생기는 점수의 수와 g 개를 만드는 동안의 총합입니다.

score_bytes 에 헤드 수를 빠뜨리기 쉽습니다. max_context_for_bytes 는 실수 제곱근을 반올림하면 한 칸 넘치는 일이 생기므로 math.isqrt 를 쓰세요 — 예산을 헤드와 바이트로 나눈 뒤 정수 제곱근을 취하면 됩니다. append_pairs(n) 은 새 질의 하나가 자기까지 포함해 n + 1 개를 보므로 n + 1 입니다. generate_pairsn+1 부터 n+g 까지의 합이고, generate_pairs(0, N)causal_pairs(N) 과 같은지 꼭 확인해 보세요.

센 숫자를 기록으로 남긴다

위 함수들을 실제로 돌려 /root/work/tf-cost/cost_report.jsond_model·d_ff·n_heads·itemsize·table·quad_ratio·linear_ratio·crossover_n·quad_at_crossover·linear_at_crossover·causal_n·causal_pairs·wasted_pairs·wasted_fraction·score_bytes_at_causal_n·max_context_1gib·append_pairs_at_causal_n·generate_pairs_at_causal_n·identity_ok 를, /root/work/tf-cost/cost_report.md## 무엇을 세었나 ## 두 배로 늘리면 무엇이 네 배가 되나 ## 교차점은 어디인가 ## 인과 마스크가 버리는 절반 ## 메모리와 한 토큰씩 늘어나는 점수 다섯 절로 기록하세요.

숫자는 손으로 적지 말고 여러분의 코드를 돌려 얻은 값으로 채우세요. 표는 길이 [128, 256, 512, 1024, 2048, 4096] 이고 quad_ratio·linear_ratio 는 마지막 두 줄 사이의 배율입니다. causal_n 은 2048, 자료형은 2바이트, 예산은 1073741824, 생성 토큰 수는 256 입니다. identity_okgenerate_pairs(0, causal_n) == causal_pairs(causal_n) 이 참인지입니다. 기록에는 재어 보지 않은 숫자를 쓰지 마세요 — 실제 모델의 초나 GB 는 여기서 잰 적이 없습니다.