본문 바로가기

SASS_Probe

attention_prob_materialization_f32 - Kernel Boundary 를 통한 Softmax Intermediate Materialization 분석

1. 실험 개요

이전 flashattention_toy_f32 실험에서 관찰된 compiler optimization 을 더 명확하게 검증하기 위해 설계

이전 실험의 materialized softmax source 에는 다음과 같은 중간 저장이 있었다.

float p = __expf(score - max_value);
p_base[j] = p;

float inv_sum = 1.0f / sum;
p_base[j] = p_base[j] * inv_sum;

소스 코드를 보면 다음 memory flow 를 예상할 수 있다.

  • STG unnormalized probability
  • LDG unnormalized probability
  • STG normalized probability

그러나 실제 SASS 에선느 중간 unnormalized probabilty 의 store / load 가 나타나지 않았다.

컴파일러는 같은 kernel 안에서 해당 값이 다시 사용된다는 사실을 추적하고, unnormalized probability 를 register 에 유지한 뒤 normalized probability 만 최종적으로 젖아

즉 다음 최적화가 발생했다.

  • source
    • p = exp(score - max)
    • store p
    • load p
    • p = p / sum
    • store p
  • optimized SASS
    • p = exp(score - max)
    • p = p / sum
    • store normalized p

이번 실험에서는 softmax 를 여러 kernel 로 분리해 중간값을 kernel boundary 밖으로 노출했다.

목표는 다음 질문을 확인하는 것이

  • 같은 softmax 계산이라도 kernel boundary 를 만들 compiler 가 제거했더 ㄴintermediate store / load SASS 의 STG  / LDG 로 다시 나타나는가

 

2. 실행 결과

attention_prob_materialization_f32 done
y[0] = -0.05316688
y[1] = -0.02316688
y[2] = 0.00683312
y[3] = -0.04475802

y[0] = -0.05316688

이전 실험의 결과와 동일하다

따라서 kernel 을 더 세분화해 intermediate tensor 를 강제로 materialize 했지만, 계산되는 attention output 은 동일하게 유지되었다.

이번 실험의 목적은 numerical result 를 바꾸는 것이 아니라 다음 차이를 만드는 것이다.

  • 동일한 수학적 계산
  • 다른 execution structure
  • 다른 memory materialization boundary

 

3. 전체 계산 구조

attention 을 네 개의 kernel 로 분리

Kernel 1
QK score 계산
    ↓
scores global memory 저장

Kernel 2
scores load
max reduction
exp 계산
sum 계산
    ↓
unnorm_probs global memory 저장
sums global memory 저장

Kernel 3
unnorm_probs load
sums load
normalize
    ↓
probs global memory 저장

Kernel 4
probs load
values load
value accumulation
    ↓
y global memory 저장

다음 flow 를 가진다

q, k
  │
  ▼
attention_score_store_kernel
  │
  ├─ STG scores
  ▼
scores
  │
  ▼
attention_exp_sum_store_kernel
  │
  ├─ LDG scores
  ├─ STG unnorm_probs
  └─ STG sums
  ▼
unnorm_probs, sums
  │
  ▼
attention_normalize_probs_kernel
  │
  ├─ LDG unnorm_probs
  ├─ LDG sums
  └─ STG probs
  ▼
probs
  │
  ▼
attention_value_from_probs_kernel
  │
  ├─ LDG probs
  ├─ LDG values
  └─ STG y

이 구조에서는 각 intermediate tensor 가 kernel boundary 를 넘어 다음 kernel 의 입력으로 사용된다.

따라서 다음 intermediate 들은 compiler 가 제거할 수 없는 외부 관찰 가능한 global memory state 가 된다.

  • scores
  • unnorm_probs
  • sums
  • probs

 

4. 분석 대상 커널

  • attention_score_store_kernel
    • QK score 계싼
  • attention_exp_sum_store_kernel
    • max, exp, sum 계산
  • attention_normalize_probs_kernel
    • probability normalization
  • attention_value_probs_kernel
    • probability-value accumulation

이번 분석에서 가장 중요한 두 kernel 은 다음이다. 

  • attention_exp_sum_store_kernel
  • attention_normalize_probs_kernel

두 kernel 사이에서  다음 memory boundary 가 만들어졌기 때문이다.

  • STG unnorm_probs
  • STG sums
  • LDG unnorm_probs
  • LDG sums

 

5. attention_score_store_kernel

5.1 역할

이 kernel 은 query row 와 각 key 사이의 dot product 를 계산한다

  • score_j = q  k_j x scale

현재 toy configuration

  • D = 4
  • N_KEYS = 4
  • scale = 0.5

따라서 한 query row 에 대해 네 개의 score 를 계산한다

5.2 SASS 구조

핵심 instruction pattern

LDG.E.CONSTANT ...
LDG.E.CONSTANT ...

FMUL.FTZ ...
FFMA.FTZ ...
FFMA.FTZ ...
FFMA.FTZ ...

FMUL.FTZ ..., 0.5

STG.E [scores]
STG.E [scores+0x4]
STG.E [scores+0x8]
STG.E [scores+0xc]

5.3 QK dot-product accumulation

각 score 계산은 첫 multiply 와 이후 fused accumulation 으로 구성된다.

수식

score =
    q0 × k0
  + q1 × k1
  + q2 × k2
  + q3 × k3
  
SASS pattern

FMUL
FFMA
FFMA
FFMA

첫 multiply

partial = q0 x k0

이후 accumulation

partial = q1 × k1 + partial
partial = q2 × k2 + partial
partial = q3 × k3 + partial

이 부분은 이전 fma_contract_f32 실험에서 확인한 multiply-add contraction 과 연결된다.

mul + add dependency
-> FFMA

 

5.4 Scale 적용

dot product 결과에 0.5 scale 이 적용된다.

FMUL / FFMA dot product

-> FMUL scale

-> STG score

 

5.5 Score materialization

마지막 네 개의 STG 는 각 score 를 global memory 에 저장한다.

STG.E [R2.64], ...
STG.E [R2.64+0x4], ...
STG.E [R2.64+0x8], ...
STG.E [R2.64+0xc], ...

이는 다음 source-level array 를 나타낸다.

scores[row, 0]
scores[row, 1]
scores[row, 2]
scores[row, 3]

따라서 첫 번째 materialization boundary 가 확인된다.

  • register QK score
    • STG
    • global scores matrix

 

6. attention_exp_sum_store_kernel

6.1 역할

이 kernel 은 저장된 scores 를 읽고 다음을 계산한다

  1. row maximum
  2. exp(score - max)
  3. sum of exponentials
  4. unnormalized probability 저장
  5. sum 저장

수식

m = max_j score_j

u_j = exp(score_j - m)

sums[row] = l

이번 실험에서 가장 중요한 kernel 이다.

 

6.2 Scores load

초반부에서 score 네 개를 global memory 로부터 읽는다.

LDG.E.CONSTANT R5,  [R2.64]
LDG.E.CONSTANT R7,  [R2.64+0x4]
LDG.E.CONSTANT R9,  [R2.64+0x8]
LDG.E.CONSTANT R11, [R2.64+0xc]

이 load 들은 이전 kernel 의 score stores 와 연결된다.

  • attention_score_store_kernel
    • STG scores
  • attention_exp_sum_store_kernel
    • LDG scores

따라서 scores matrix 가 kernel 사이에서 global memroy tensor 로 존재한다.

 

6.3 Maximum reduction

다음 SASS 가 등장

FMNMX.FTZ R6, R5, -3.40282346638528859812e+38, !PT
FMNMX.FTZ R6, R6, R7, !PT
FMNMX.FTZ R6, R6, R9, !PT
FMNMX.FTZ R6, R6, R11, !PT

이는 다음 reduction 에 대응한다.

m = max(-∞, s0)
m = max(m, s1)
m = max(m, s2)
m = max(m, s3)

즉 max reduction 은 branch 가 아니라 FMNMX chain 으로 구현된다.

  • source
    • fmaxf
  • SASS
    • FMNMX

 

6.4 Stable softmax subtraction

각 score 에서 maximum 을 뺀다

FADD.FTZ R8, R7, -R6
FADD.FTZ R5, R5, -R6
FADD.FTZ R9, R9, -R6
FADD.FTZ R6, R11, -R6

SASS 에서는 subtraction 도 FADD 와 negated operand 로 표현된다.

  • score - max
    • FADD score, -max

이 단계는 numerical stability 를 위한 stable softmax 구조다

exp(score)

대신

exp(score - max(score))

를 계산한다.

 

6.5 Exponential calculation

각 값은 먼저 log2(e) 를 곱한다

FMUl.FTZ ..., 1.444...

그 다음

MUFU.EX2

가 실행된다.

이는 다음 관계 때문이다.

exp(x) = 2^(x x log2(e))

따라서 __expf 의 SASS signature 는 다음과 같다

  • FADD score - max
  • FMUL log2(e)
  • MUFU.EX2

이번 kernel 에서 네 개의 exponential 이 각각 MUFU.EX2 로 나타났다.

 

6.6 Sum reduction

exponential 결과를 더한다.

관찰된 SASS

FADD.FTZ R10, R7, R8
FADD.FTZ R10, R9, R10
FADD.FTZ R13, R11, R10

수식

  • sum = u0 + u1 + u2 + u3

이번 toy size 에서는 sequential FADD chain 으로 reduction 된다.

 

6.7 Unnormalized probability materialization

가장 중요한 부분

STG.E [R2.64], R7
STG.E [R2.64+0x4], R8
STG.E [R2.64+0x8], R9
STG.E [R2.64+0xc], R11

다음을 저장

unnorm_probs[row, 0]
unnorm_probs[row, 1]
unnorm_probs[row, 2]
unnorm_probs[row, 3]

이전 flashattnetion_toy_f32 에서는 unnormalized probability 가 같은 kernel 안에서 normalize 되었기 때문에 compiler 가 중간 store 를 제거했다.

이번에는 다음 kernel 에서 이 값을 소비하므로 제거할 수 없다.

  • attention_exp_sum_store_kernel
    • STG unnorm_probs
  • attention_normalize_probs_kernel
    • LDG unnorm_probs

따라서 source-level intermediate tensor 가 실제 SASS-level global memory tensor 로 materialize 되었다.

 

6.8 Sum materialization

마지막 store

STG.E [R4.64], R13

는 row sum 을 저장한다.

sum[row] = u0 + u1 + u2 + u3

따라서 이 kernel 에는 두 종류의 output materialization 이 있다.

  • vector output
    • unnorm_probs[row, :]
  • scaler output
    • sums[row]
  • SASS signature
    • 4 x STG unnorm_probs
    • 1 x STG sum

 

7. attention_normalize_probs_kernel

7.1 역할

이  kernel 은 이전 kernel 이 저장한 unnormalized probability 와 row sum 을 읽는다.

수식

  • inv_sum = 1 / sums[row]
  • probs[row, j] = unnorm_probs[row, j] x inv_sum

 

7.2 Sum load

다음 load 가 나타난다.

LDG.E.CONSTANT R6, [R6.64]

이는

sum = sums[row]

에 대응한다.

즉 이전 kernel 의 

STG sums 

가 현재 kernel 의

LDG sums

로 이어진다.

 

7.3 Unnormalized probability load

네 개의 probability 를 global memory 에서 읽는다.

LDG.E.CONSTANT R9,  [R2.64]
LDG.E.CONSTANT R11, [R2.64+0x4]
LDG.E.CONSTANT R13, [R2.64+0x8]
LDG.E.CONSTANT R15, [R2.64+0xc]


이는 다음을 의미한다.

u0 = unnorm_probs[row, 0]
u1 = unnorm_probs[row, 1]
u2 = unnorm_probs[row, 2]
u3 = unnorm_probs[row, 3]

이 결과로 이번 시렇ㅁ의 핵심 materialization chian 이 완성된다.

  • kernel 2
    • STG unnorm_probs
  • kernel 3
    • LDG unnorm_probs

 

7.4 Reciprocal

row sum 의 reciprocal 은

MUFU.RCP R8, R6

로 계산된다.

수식

inv_sum = 1 / sum

--use_fast_math 환경이므로 reciprocal 은 special function unit 의 MUFU.RCP 로 내려간다.

 

7.5 Normalize

각 unnormalized probability 에 reciproccal 을 곱한다

FMUL.FTZ R9,  R8, R9
FMUL.FTZ R11, R8, R11
FMUL.FTZ R13, R8, R13
FMUL.FTZ R15, R8, R15


수식:


p0 = u0 × inv_sum
p1 = u1 × inv_sum
p2 = u2 × inv_sum
p3 = u3 × inv_sum

 

7.6 Normalized probability materialization

계산된 probabilites 를 global memory 에 저장한다.

STG.E [R4.64], R9
STG.E [R4.64+0x4], R11
STG.E [R4.64+0x8], R13
STG.E [R4.64+0xc], R15


따라서 다음 chain이 확인된다.


STG unnorm_probs
→ kernel boundary
→ LDG unnorm_probs
→ MUFU.RCP
→ FMUL normalize
→ STG probs

이것이 이번 실험에서 보고자 했던 핵심 구조다

 

8. attention_value_from_probs_kernel

8.1 역할

이 kernel 은 normalized probability 와 values 를 읽어 최종 output 을 계산한다

  • 수식
    • y = SIGMA_j probs_j x v_j
  • 각 output component
    • y_d = SIGMA_j probs_j x v[j, d]

 

8.2 Input loads

초반부에서 다음 데이터들을 읽는다.

probs[row, 0:4]

values[0:4, 0:4]

SASS 에서는 여러 LDG.E.CONSTANT로 나타난다.

이전 kernel 의 

STG probs

가 현재 kernel 의 

LDG probs

로 이어진다.

 

8.3 Value accumulation

핵심 arithmetic 은 FFMA chain 으로 나타난다

FFMA.FTZ R0, R0, R3, RZ
FFMA.FTZ R0, R15, R17, R0

수식

acc = p x v + acc

첫 accumulation 에서도 compiler 는 zero addend 를 이용해 FFMA 를 사용한다.

acc0 = p0 x v0 + 0

이후

acc0 = p1 x v1 + acc0
acc0 = p2 x v2 + acc0
acc0 = p3 x v3 + acc0

로 누적된다.

즉 probaility-value accumulation 의 SASS signature 는 다음이다.

LDG probs

LDG values

FFMA chain

 

8.4 Final output materialization

마지막에는 output vector 를 저장한다

STG.E [R2.64], R23
STG.E [R2.64+0x4], R9
STG.E [R2.64+0x8], R11
STG.E [R2.64+0xc], R13

이는

y[row, 0:4]

의 최종 output store 다.

 

9. 전체 Materialization Chain

이번 실험에서 확인된 전체 global memory flow 는 다음과 같다

Q/K
 │
 │ LDG
 ▼
QK score registers
 │
 │ STG
 ▼
scores
 │
 │ LDG
 ▼
max / exp / sum registers
 │
 ├─ STG unnorm_probs
 └─ STG sums
        │
        │ LDG
        ▼
normalize registers
 │
 │ STG
 ▼
probs
 │
 │ LDG
 ▼
value accumulator registers
 │
 │ STG
 ▼
y


instruction 관점에서 축약하면:


LDG q/k
FMUL/FFMA
STG scores

LDG scores
FMNMX
FADD
FMUL
MUFU.EX2
FADD
STG unnorm_probs
STG sums

LDG unnorm_probs
LDG sums
MUFU.RCP
FMUL
STG probs

LDG probs
LDG values
FFMA
STG y

 

10. 이전 실험과의 비교

10.1 flashattention_toy_f32 의 same-kernel softmax

이전 materialized softmax kernel 의 source 에는 중간 store 가 있었다.

  • exp result
    • p_bas store
    • p_base load
    • normalize
    • final store

그러나 실제 SASS 에서는 compiler 가 중간 store/laod 를 제거

실제 구조

  • exp result register
    • normalize register
    • final STG probs

이유

  • 같은 kernel 내부
  • 외부 관찰 가능성 없음
  • store 후 바로 동일 kernel 에서 소비
  • compiler 가 data flow 전체를 볼 수 있음

 

10.2 이번 kernel-boundary 실험

이번 실험에서는 분리

  • kernel A
    • STG unnorm_probs
    • STG sums
  • kernel B
    • LDG unnorm_probs
    • LDG sums

컴파일러는 별도 kernel lauch 사이에서 register value 를 직접 전달할 수 없다.

각 kernel 은 독립된 global function 이고 intermediate result 는 global memory 를 통해 전달되어야 한다.

따라서 store/load 가 제거되지  않았다.

 

11. 핵심 연구 결론

이번 실험은 다음 사실을 명확하게 보여준다.

Materialization 은 source code 에 배열 대입문이 있다는 사실만으로 결정되지 않는다.

같은 kernel 안에서 intermediate store / load 가 외부에 관찰되지 않으면 compiler 는 이를 제거할 수 있다. 

  • source STG / LDG intent
    • optimized register value

반면 kernel boundary 를 넘는 값은 다음 kernel 이 소비할 수 있도록 global memory 로 존재해야 한다.

  • producer kernel
    • STG global intermediate
    • kernel boundary
    • LDG global intermediate
    • consumer kernel

따라서 진짜 materialization boundary 는 다음과 같은 외부 관찰 가능성에 의해 만들어진다.

  • kernel boundary
  • volatile access
  • function / call boundary
  • external consumer
  • synchronization boundary
  • aliasing or observale memory semantics

 

12. FMA 실험과의 연결

source temporary variable 는 반드시 SASS intermediate value 로 남지 않는다.

tmp 는 독립적인 FMUL 결과로 materialize 되지 않았다.

이번 실험에서는 같은 개념을 tensor 수준으로 확장

 

13. Flash Attention 과의 연결

일반적인 amterialized attention 은 다음 중간 행렬을 만든다

  • S = QK
  • P = softmax(S)
  • O = PV

이번 실험은 여기에 unnormalized probability 와 sum 까지 추가로 노출

반면 online attention 은 당므 중간값을 register 상태로 유지

  • score
  • running max
  • running sum
  • value accumulator

그래서 SASS 에서 global intermediate memory boundaries 가 사라지낟.

최종 output 만 저장

 

14. SASS 에서 Materialization 을 판별하는 방법

14.1 단순 STG 개수만 세면 안 된다

모든 kernel 에는 output store 가 있을 수 있다.

온라인 attention 에도 최종 y 저장을 위한 STG 는 존재한다.

정확하게

어떤 logical tensor 를 저장하는 STG 인가에 대한 질문

 

14.2 Producer-consumer pair 확인

materializatoin 을 확정하려면 producer kernel 과 consumer kernel 을 함께 본다.

  • producer
    • STG unnorm_probs
  • Consumer
    • LDG unnorm_probs

이 pair 가 존재하면 intermediate tnesor 가 실제 global memory 를 통해 전달된다고 볼 수 있다. 

 

14.3 주소와 kernel parameter 연결

SASS 만 보면 register 주소가 어떤 source pointer 인지 바로 드러나지 않을 수 있다. 

따라서 다음을 같이 봐야 한다

  • kernel parameter order
  • IMAD.WIDE address calculation
  • LDG/STG offset pattern
  • source array shape
  • consumer kernel 의 대응 load

예를 들어 네 개의 연속 store

  • base
  • base + 0x4
  • base + 0x8
  • base + 0xc

는 연속된 네 개 FP32 element 를 나타낸다.

 

14.4 최종 판별 프레임

  1. source-level intermediate 를 식별
  2. producer kernel 의 STG 를 찾는다
  3. consumer kernel 의 LDG 를 찾는다
  4. 두 accuess 가 같은 logical tensor 인지 확인
  5. kernel boundary 가 존재하는지 확인
  6. intermediate 가 register 로 제거되었는지,  memory 에 남았는지 판단한다.

 

가장 중요한 핵심

Attention 최적화의 핵심은 산술 연산을 줄이는 것만이 아니라, 중간 tensor 가 global memory state 로 materialization 되는 경계를 제거하는 것이다.