1. 실험 개요
이번 실험은 attention 의 중간 score 가 다음 세 가지 위치 중 어디에 존재하는지 SASS 수준에서 구분하기 위한 실험이다.
- register-resident
- shared-memory-resident
- global-memory-resident
앞선 실험에서는 kernel boundary 를 기준으로 intermediate tensor 가 global memory 에 materialize 되는 구조를 확인했다.
- producer kernel
- STG intermediate
- kernel boundary
- LDG intermediate
- consumer kernel
kernel 을 분리하지 않고, 하나의 kernel 내부에서 서로 다른 thread 또는 warp 가 계산 결과를 공유하도록 구성
핵심 질문은 다음과 같다
- Attention score 를 같은 thread 가 바로 소비하면 register 에 유지되는가
- Attention score 를 다른 warp 가 소비하도록 만들면 shared memory 에 실제로 materialize 되는가?
- 이 차이가 SASS 에서 STS, BAR, LDS instruction 으로 나타나는가
이를 위해 하나의 CUDA 파일 안에 두 가지 구현을 만들었다.
- attention_register_reference_kernel
- attention_shared_scores_kernel
두 kernel 은 동일한 attention 결과를 계산하지만 score 의 전달 방식이 다르다.
2. 실험 구성
Toy attention configuration
- D = 4
- N_KEYS = 4
- DV = 4
- scale = 0.5
계산식
- score_j = dot(q, k_j) x scale
- p_j = softmax(score_j)
- y = SIGMA_j p_j x v_j
두 kernel 은 같은 수식을 사용
차이는 score intermediate 를 전달하는 방식
Register reference
- 하나의 thread 가 score 4 개를 모두 계산
- score 를 register 에 유지
- 같은 thread 가 softmax 수행
- 같으 thread 가 value accumulation 수행
Shared score path
- warp 0 의 lane 0 이 score 0 계산
- warp 1 의 lane 0 이 score 2 계산
- warp 2 의 lane 0 이 score 2 계산
- warp 3 의 lane 0 이 score 3 계산
각 warp 가 score 를 shared memory 에 저장
- block-wide barrier
- thread 0 이 score 4 개를 shared memory 에서 읽음
- softmax 와 value accumulation 수행
3. 실행 결과
attention_score_shared_f32 done
reference y[0:4] = -0.05316688 -0.02316688 0.00683312 -0.04475802
shared y[0:4] = -0.05316688 -0.02316688 0.00683312 -0.04475802
max_abs_diff = 0.00000000e+00
두 구현의 출력은 완전히 동잃
따라서 두 kernel 은 동일한 수학적 attention 을 계산
이 실험에서 바뀐 것은 numerical semantics 가 아니라 execution struction
- Register reference
- producer 와 consumer 가 같은 thread
- Shared score
- producer 와 consumer 가 서로 다른 warp/thread
4. attention_register_regerence_kernel
4.1 역할
한 thread 가 한 query row 를 전부 담당한다
- QK score 4개 계산
- row maximum
- exponentials
- denominator
- normalized probabilities
- value accumulation
- output 저장
score intermediatino 를 다른 thread 에 전달하지 않는다.
따라서 scrore 는 register 에 유지될 수 있다.
4.2 Q / K / V global load
입력 q, k, v 는 global memory 에서 읽힌다.
SASS 에서는 다음 계열로 나타난다.
LDG.E.CONSTANT
여기서 LDG.E.CONSTANT 의 CONSTANT 를 CUDA constant memory 사용으로 바로 해석할 필요는 없다.
이번 분석에서 중요한 의미는 다음이다.
- 입력 q / k / v
- globa 또는 read-only memory load
4.3 QK dot product
score 하나의 계산은 다음 형태로 나타난다.
- FMUL.FTZ
- FMUL.FTZ
- FMUL.FTZ
- FMUL.FTZ
- FMUL.FTZ ... , 0.5
수식
- score = q0 x k0 + q1 x k1 + q2 x k2 + q3 x k3
- score = score x 0.5
첫 번째 multiplication
- FMUL
이후 accumulation
- FFMA
- FFMA
- FFMA
마지막 scale
- FMUL 0.5
따라서 score caculation signature 는 다음과 같다
- FMUL
- FFMA chain
- FMUL scale
4.4 Register-resident score
register reference kernel 에서는 score 계산 이후 다음 instruction 으로 바로 연결
- score arithmetic
- FMNMX
- FADD
- MUFU.EX2
중간에 다음 instruction 이 없다.
- STS
- LDS
- STG scores
- LDG scores
- BAR.SYNC
이는 score 가 별도의 memory object 로 materialize 되지 않았음을 의미한다.
- score calculation result
- register
- softmax consumer
즉 producer 와 consumer 사이에 memory boundary 가 없다.
4.5 Softmax maximum
row maximum 은 FMNMX chain 으로 계산된다
- FMNMX.FTZ
- FMNMX.FTZ
- FMNMX.FTZ
- FMNMX.FTZ
수식
- m = max(-INF, s0)
- m = max(m, s1)
- m = max(m, s2)
- m = max(m, s3)
source 의 fmaxf 가 SASS 에서 branch 없이 FMNMX 로 내려간다.
4.6 Stable softmax
각 score 에서 maximum 을 뺀다
score_j - max_score
SASS 에서는 negated operand 를 사용하는 FADD 로 나타난다.
FADD.FTZ score, score, -max
이후 __expf 는 다음 패언을 나타난다.
FMUL.FTZ ..., 1.442...
MUFU.EX2 ...
1.442.. 는 log2(e) 다
exp(x) = 2^(x x log2(e))
따라서 fast exponential signature 는 다음과 같다
- FADD score - max
- FMUL log2(e)
- MUFU.EX2
4.7 Normalize
exponential sum 은 FADD chain 으로 계산된다.
- sum = p0 + p1 + p2 + p3
reciprocal
- MUFU.RCP
normalizatoin
- FMUL.FTZ
즉
- inv_sum = 1 / sum
- prob_j = exp_j x inv_sum
이다.
probability 도 memory 에 저장되지 않고 register 에 유지된다.
4.8 Value accumulation
probability 와 value 의 multiplication / accumulation 은 FMUL 과 FFMA 로 나타난다.
- acc_d = p0 x v0,d
- acc_d = p1 x v1,d + acc_d
- acc_d = p2 x v2,d + acc_d
- acc_d = p3 x v3,d + acc_d
SASS signature
- FMUL
- FFMA
- FFMA
- FFMA
마지막 output 만 global memory 에 저장한다.
- STG.E [output]
- STG.E [output+0x4]
- STG.E [output+0x8]
- STG.E [output+0xc]
4.9 Register reference 요약
- Q/K/V
- LDG
- scores
- register-resident
- probability
- register-resident
- accumulators
- register-resident
- output
- STG
- SASS signature
- LDG
- FMUL / FFMA
- FMNMX
- FADD
- MUFU.EX2
- MUFU.RCP
- FMUL / FFMA
- STG
나타나지 않는 것
- STS
- LDS
- BAR.SYNC
5. attention_shared_scores_kernel
5.1 역할
이 kernel 은 하나의 block이 하나의 query row 를 처리한다.
- block size
- 128 threads
즉 4 개의 warp 가 존재한다.
- warp 0
- thread 0~31
- warp 1
- thread 32~63
- warp 2
- thread 64~95
- awrp 3
- thread 96~127
각 warp 의 lane 0 만 score 를 계산한다.
- thread 0
- score 0
- thead 32
- score 1
- thread 64
- score 2
- thread 96
- score 3
score 는 shared memory 에 저장된다.
shared_scores[warp_id] = score;
후속 softmax 는 thread 0 이 담당한다
5.2 Thread ID 읽기
- SASS
- S2R R6, SR_TID.x
이는 threadIdx.x 를 읽는 instruction 이다.
5.3 Lane ID 와 warp ID
- source
- lane_id = threadIdx.x & 31;
- warp_id = threadIdx.x >> 5;
lane 조건은 bitwise instruction 과 predicate 로 구현된다.
warp ID 계산은 다음과 같다
SHF.R.S32.HI R6, RZ, 0x5, R6
이는 의미상
threadIdx.x >> 5
즉 32 로 나누어 warp ID 를 구한다.
5.4 Conditional producer execution
- source 조건
- if (lane_id == 0 && warp_id < 4)
따라서 네 warp 의 lane 0 만 QK score 를 계산한ㄷ
나머지 thread 는 score 계산과 shared store 를 건너뛴다.
SASS 에는 predicate setup 과 branch 가 나타난다.
- predicate calculation
- conditional branch
- producer path
5.5 QK score calculation
각 producer thread 가 query 와 자신의 key 를 읽는다.
LDG.E.CONSTANT
dot product
- FMUL.FTZ R7, R7, R8
- FMUL.FTZ R7, R10, R9, R7
- FMUL.FTZ R7, R12, R11, R7
- FMUL.FTZ R7, R14, R13, R7
- FMUL.FTZ R7, R7, 0.5
이는 정확히 다름 수식
- score = q0 x k0 + q1 x k1 + q2 x k2 + q3 x k3
- score - score x 0.5
산술 구조는 register reference 와 동일
6. Shared-memory materialization
6.1 Shared store
이번 실험의 첫 번째 핵심 instruction
- STS [R6.X4], R7
source
- shared_scores[warp_id] = score
R6 에는 warp ID 가 들어 있다.
.X4 는 index 를 4-byte 단위로 확장하는 주소 표현이다.
따라서 개념적 주소는 다음과 같다
- warp 0
- shared_scores[0]
- warp 1
- shared_scores[1]
- warp 2
- shared_scores[2]
- warp 3
- shared_scores[3]
score 가 register 에서 shared memory 로 이동했다.
- register score
- STS
- shared memory score
이것이 shared-memory materialization 의 producer side 다.
6.2 Control-flow reconvergence
producer 조건 때문에 일부 thread 만 QK 계산과 shared store 를 실행한다.
SASS 에는 다음 instruction 이 나타난다.
- BSSY B0,
- BSYNC B0
이 instruction 은 block-wide memory barrier 가 아니라 divergent control flow 의 reconvergence 를 관리한다.
- BSSY
- synchronization point 설정
- BSYNC
- 분기 경로 reconvergence
이를 __syncthreads() 에 직접 대응시키면 안 된다.
6.3 Block-Wide barrier
실제 __syncthreads() 에 대응하는 instruction 은 다음이다.
- BAR.SYNC.DEFER_BLOCKING 0x0
이 barrier 는 네 warp 가 shared memory store 를 완료할 때까지 block전체를 동기화한다.
필요 이유
- score producer
- warp 0~3
- score consumer
- thread 0
warp 0 의 thread 0 이 shared memory 를 읽는 시점에 warp 1~3 의 store 가 완료됐다는 보장은 없다.
따라서 다음 ordering 이 필요하다
- STS score 0~3
- BAR.SYNC
- LDS score 0~3
6.3 Barrier 와 early exit
source 에는 barrier 이후 다음 조건이 있다.
- fi(threadIdx.x != 0 ) { return; }
SASS 에서도 barrier 뒤에 predicated exit 가 나타난다
- @p1 EXIT
즉 exectuion order 는 다음과 같다
- 전체 thread
- producer branch 수행 또는 skip
- 전체 thread
- BAR.SYNC 도달
- thread 1~127
- EXIT
- thread 0
- shared load 및 후속 softmax 수행
모든 thread 가 barrier 를 통과한 당므 일부 thread 가 종료되므로 올바른 synchronization 구조
7. Shared-memory load
7.1 Source-level loads
네 개의 scalar load 존재
- const float s0 = shared_scores[0];
7.2 Vectorized LDS
실제 sASS 에서는 하나의 instruction 으로 합쳐졌다.
- LDS.128 R4, [RZ]
128 bit 는 다음과 같다
- 128 bit
- 16 byte
- 4 x FP32
따라서 개념적으로
- R4 = shared_scores[0]
- 5
- 6
- 7
가 된다.
이것은 compiler 가 연속성과 alignment 를 이용해 scaler load 네 개를 vectorized shared load 하라ㅗ 합친 결과다.
- source
- 4 x scaler shared load
- SASS
- 1 x LDS.128
따라서 source-level memory access 개수와 SASS instruction 개수는 반드시 같지 않다.
7.3 Shared producer-consumer chain
이번 실험의 핵심 chain
- producer registers
- STS
- shared_scores
- BAR.SYNC
- LDS.128
- consumer registers
- SASS
- STS [R6.X4], R7
- BAR.SYNC.DEFER_BLOCKING 0x0
- LDS.128 R4, [RZ]
이 세 instruction 이 shared-memory materialization 을 확정한다.
8. Shared load 이후 softmax
8.1 Maximum reduction
LDS.128 결과가 들어 있는 R4~R7 에 대해 max reduction 을 수행한다.
- FMNMX.FTZ
- FMNMX.FTZ
- FMNMX.FTZ
- FMNMX.FTZ ...
수식
- m = max(s0, s1, s2, s3)
8.2 Stable subtraction
- FADD.FTZ
- FADD.FTZ
- FADD.FTZ
- FADD.FTZ
수식
- s_j = score_j - max_score
8.3 Exponential
- FMUL.FTZ
- MUFU.EX2
네 score 모두 fast exponential 경로를 사용한다.
- exp(x)
- x x log2(e)
- 2^x
8.4 Sum 과 reciprocla
- FADD.FTZ
- FADD.FTZ
- FADD.FTZ
- MUFU.RCP
수식
- sum = p0 + p1 + p2 + p3
- inv_sum = 1 / sum
8.5 Probability residency
정규화된 probabilities 는 memory 에 저장되지 않는다.
- no STS probs
- no LDS probs
- no STG probs
- no LDG probs
softmax probabilities 는 register 에서 value accumulation 으로 바로 전달된다
- shared scores
- LDS.128
- register softmax
- register probabilities
- register value accumulation
따라서 이 kernel 에는 서로 다른 두 residency 가 공존한다.
- scores
- shared-memory-resident
- probabilities
- register-resident
9. Value accumulation
values 는 global memory 에서 읽힌다.
- LDG.E.CONSTANT
probability 와 value accumulation 은 FMUL 및 FFMA 로 처리된다.
최종 output 은 global memory 에 저장된다.
10. 연구 결론
실행 결과
Register reference 와 shared path 의 output 이 완전히 동일
SASS 결과
- Register reference
- shared-memory instruction 없음
- Shared path
- STS
- BAR.SYNC
- LDS.128
따라서 attentnion intermediate 의 residency 를 SASS 수준에서 다음처럼 구분할 수 있다.
- Register-resident
- producer result 가 register 로 consumer 에 직접 연결
- Shared-memory-resident
- producer result 가 STS 로 저장되고
- barrier 이후 LDS 로 다시 읽힘
- Global-memory-resident
- producer result 가 STS 로 저장되고
- barrier 이후 LDS 로 다시 읽힘
- Global-memory-resident
- producer kernel 이 STG 로 저장하고
- 다음 kernel 이 LDG 로 다시 읽음
핵심 문장
중간 tensor 의 materialization 여부 뿐 ㅁ아니라 어느 memory hierarchy 에 materialize 되는지도 SASS 의 instruction pattern으로 구분할 수 잇다.
11. Attention 최적화 관점의 의미
Attention optimization의 핵심은 단순히 arithmetic instruction 수를 줄이는 데 있지 않다.
같은 QK, softmax, PV 계산을 수행하더라도 intermediate가 어디에 존재하는지에 따라 execution cost가 달라진다.
Register:
가장 가까운 storage
thread-local
explicit synchronization 없음
Shared memory:
block-local
thread/warp 간 communication 가능
explicit synchronization 필요
Global memory:
device-wide
kernel 간 communication 가능
높은 memory traffic과 kernel boundary 발생
따라서 실제 optimized attention kernel을 분석할 때는 다음을 함께 봐야 한다.
산술 primitive:
FFMA
FMNMX
MUFU.EX2
MUFU.RCP
memory residency:
register
shared
global
communication:
same-thread dependency
warp shuffle
shared memory
global memory
synchronization:
none
warp synchronization
block barrier
kernel boundary