본문 바로가기

SASS_Probe

attention_score_shared_f32 - Register 와 Shared Memory 사이의 Attention Score Materialization 분석

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