본문 바로가기

SASS_Probe

online_softmax_f32 분석

좋습니다. 이 online_softmax_f32 결과는 softmax_small과 완전히 다른 구조입니다.

핵심 결론부터 말하면:

online_softmax_f32는 shared memory를 쓰지 않는다.
running max와 running sum을 register accumulator로 유지한다.

이전 softmax_small_f32에서는 STS, LDS, BAR.SYNC가 계속 나왔습니다. 그런데 이번 online_softmax_f32 SASS에서는 그런 shared memory reduction 패턴이 없습니다. 대신 대량의 LDG, FSETP/FSEL, FADD, FMUL, MUFU.EX2, FFMA, MUFU.RCP, STG가 반복됩니다. 즉 중간 exp 벡터를 shared memory에 저장하지 않고, register 상에서 online update를 수행하는 방향으로 내려갔습니다.


online_softmax_f32 분석

문서 위치 추천:

notes/05_nn_ops/online_softmax_f32.md

1. CUDA 코드 구조

우리가 작성한 커널은 대략 이런 구조였습니다.

__global__ void online_softmax_f32_kernel(const float* __restrict__ x,
                                          float* __restrict__ y,
                                          int n) {
    if (threadIdx.x == 0 && blockIdx.x == 0) {
        float m = -FLT_MAX;
        float s = 0.0f;

        for (int i = 0; i < n; ++i) {
            float v = x[i];

            float new_m = m > v ? m : v;

            float old_scale = expf(m - new_m);
            float new_term = expf(v - new_m);

            s = s * old_scale + new_term;
            m = new_m;
        }

        for (int i = 0; i < n; ++i) {
            y[i] = expf(x[i] - m) / s;
        }
    }
}

high-level 의미:

1. 전체 x를 한 번 훑으며 running max m과 running sum s를 갱신한다.
2. 최종 m, s를 이용해 y[i] = exp(x[i] - m) / s를 저장한다.

2. 가장 먼저 보이는 구조: 단일 thread guard

SASS 초반:

/*0010*/ S2UR UR4, SR_CTAID.X ;
/*0020*/ S2R R0, SR_TID.X ;
/*0030*/ LOP3.LUT P0, RZ, R0, UR4, RZ, 0xfc, !PT ;
/*0040*/ @P0 EXIT ;

이건 CUDA의 이 조건에 대응됩니다.

if (threadIdx.x == 0 && blockIdx.x == 0) {
    ...
}

즉 threadIdx.x != 0이거나 blockIdx.x != 0이면 바로 종료합니다.

이 커널은 병렬 softmax가 아니라 SASS 패턴 관찰용 단일 thread online softmax입니다. 그래서 shared memory도, barrier도, block reduction도 없습니다.


3. 초기 accumulator

초반에 다음 값들이 보입니다.

/*0070*/ MOV R21, RZ ;
/*0090*/ MOV R11, 0xff7fffff ;

의미는 대략:

float s = 0.0f;
float m = -FLT_MAX;

로 볼 수 있습니다.

R11 = running max m
R21 = running sum s

이후 긴 루프 전체에서 R11은 max accumulator, R21은 sum accumulator 역할로 계속 등장합니다.

이게 softmax_small과의 가장 큰 차이입니다.

softmax_small:
    max partial, exp, sum partial이 shared memory에 저장됨

online_softmax:
    m, s가 register accumulator로 유지됨

4. online max update 패턴

중간에 반복적으로 이런 패턴이 나옵니다.

FSETP.GT.FTZ.AND P2, PT, R11, R18, PT ;
FSEL R19, R11, R18, P2 ;

의미:

new_m = (m > v) ? m : v;

즉:

new_m = max(m, v);

입니다.

여기서 R11은 기존 running max, R18 같은 레지스터는 새로 load한 x[i] 값입니다. FSETP + FSEL로 max를 만들고 있습니다.

이전 relu_f32에서는 FMNMX가 나왔지만, 여기서는 여러 값에 대한 max update가 unrolled된 상태라 FSETP + FSEL 패턴이 반복됩니다.


5. exp 변환 패턴

반복적으로 다음 3단계가 나옵니다.

FADD.FTZ ...
FMUL.FTZ ..., 1.4426950216293334961 ;
MUFU.EX2 ...

의미는 이전 softmax_small과 같습니다.

expf(a) 
→ exp2(a * log2(e))

즉:

FADD    : m - new_m 또는 v - new_m
FMUL    : * log2(e)
MUFU.EX2: exp2 계산

예를 들어 online update의 두 항은 다음입니다.

old_scale = expf(m - new_m);
new_term  = expf(v - new_m);

SASS에서는 각각:

m - new_m
→ * 1.442695...
→ MUFU.EX2

v - new_m
→ * 1.442695...
→ MUFU.EX2

형태로 나타납니다.


6. 핵심: running sum update가 FFMA로 내려감

가장 중요한 패턴은 이겁니다.

FFMA.FTZ R21, R6, R21, R7 ;

또는 중간 unrolled 구간에서 보이는 여러 FFMA.FTZ들:

FFMA.FTZ R29, R29, R21, R28 ;
FFMA.FTZ R34, R29, R25, R26 ;
FFMA.FTZ R12, R12, R20, R17 ;
...

의미는 기본적으로 이런 형태입니다.

s = s * old_scale + new_term;

즉 online softmax의 핵심 update:

s = s * expf(old_m - new_m) + expf(v - new_m);

가 SASS에서 FFMA로 내려간 것입니다.

이건 매우 중요합니다.

이전에 fma_f32에서 봤던:

a * b + c → FFMA

패턴이 여기서 실제 algorithmic accumulator update로 나타난 겁니다.

정리하면:

source:
    s = s * old_scale + new_term

SASS:
    FFMA.FTZ s_new, old_scale, s_old, new_term

즉 online softmax의 running sum update는 SASS에서 FMA accumulator update로 보입니다.


7. shared memory materialization 없음

softmax_small_f32에서는 이런 패턴이 있었습니다.

STS
BAR.SYNC
LDS
LDS
FADD/FSEL
STS
BAR.SYNC

이번 online_softmax_f32에서는 이런 shared memory 기반 reduction 패턴이 없습니다.

특히 다음이 보이지 않는 점이 중요합니다.

STS [tid...]
LDS [...]
BAR.SYNC

대신 흐름은 이렇게 됩니다.

LDG x[i]
FSETP/FSEL로 new_m 계산
FADD/FMUL/MUFU.EX2로 scale 계산
FFMA로 s update
다음 i로 진행

즉 중간값은 대부분 register에 머뭅니다.

m     : register accumulator
s     : register accumulator
exp들 : register temporary

이게 이번 실험의 핵심입니다.


8. compiler unrolling 관찰

SASS를 보면 LDG.E.CONSTANT가 16개씩 연속으로 나오는 구간이 있습니다.

LDG.E.CONSTANT R18, [R2.64] ;
LDG.E.CONSTANT R20, [R2.64+0x4] ;
LDG.E.CONSTANT R15, [R2.64+0x8] ;
...
LDG.E.CONSTANT R8, [R2.64+0x3c] ;

이건 한 번에 16개의 float 값을 load하는 unrolled loop입니다.

즉 source에는:

for (int i = 0; i < n; ++i)

였지만, SASS에서는 컴파일러가 여러 원소를 한 묶음으로 펼쳤습니다.

1개씩 처리하는 for loop
→ 16개 단위 unrolled block
→ 8개 단위 tail
→ 4개 단위 tail
→ 1개 단위 tail

실제로 뒤쪽에는 8개, 4개, 1개 tail 처리에 해당하는 유사 패턴들도 보입니다.

이것도 중요한 결론입니다.

SASS에서는 source-level loop 구조가 그대로 보존되지 않는다.
컴파일러가 unrolling과 tail path를 만들어낸다.

9. normalization phase

online update가 끝난 뒤에는 최종 m = R11, s = R21을 사용해서 output을 씁니다.

핵심 패턴:

MUFU.RCP R7, R21 ;
...
FADD.FTZ R16, R16, -R11 ;
FMUL.FTZ R16, R16, 1.4426950216293334961 ;
MUFU.EX2 R16, R16 ;
FMUL.FTZ R9, R16, R7 ;
STG.E [R4.64], R9 ;

의미:

inv_s = 1.0f / s;
e = expf(x[i] - m);
y[i] = e * inv_s;

즉 division도 역시:

e / s

가 아니라 SASS에서는:

e * rcp(s)

입니다.

MUFU.RCP
FMUL

형태로 내려갑니다.


10. softmax_small vs online_softmax 비교

softmax_small_f32

x
→ shared memory
→ max reduction
→ shared memory
→ exp(x - max)
→ shared memory
→ sum reduction
→ shared memory
→ normalize
→ y

SASS 핵심:

STS
BAR.SYNC
LDS
FSETP/FSEL
STS
BAR.SYNC
MUFU.EX2
STS
LDS
FADD
STS
MUFU.RCP
FMUL
STG

online_softmax_f32

x stream
→ register running max m
→ register running sum s
→ final normalize
→ y

SASS 핵심:

LDG
FSETP/FSEL
FADD
FMUL
MUFU.EX2
FFMA
MUFU.RCP
FMUL
STG

차이는 명확합니다.

softmax_small:
    중간 exp 값과 reduction partial이 shared memory에 materialize됨

online_softmax:
    running max/sum update가 register accumulator로 진행됨

11. 분석 프레임 적용

CUDA source expression

new_m = max(m, v);
old_scale = expf(m - new_m);
new_term = expf(v - new_m);
s = s * old_scale + new_term;
m = new_m;

compiler graph rewrite

max:
    FSETP + FSEL

exp:
    subtract
    multiply by log2(e)
    MUFU.EX2

sum update:
    fused multiply-add

SASS primitive

FSETP
FSEL
FADD
FMUL
MUFU.EX2
FFMA

intermediate materialization 여부

shared/global memory materialization 없음.

m, s는 register accumulator로 유지된다.
exp intermediate는 temporary register로 생성되고,
s update에 바로 사용된다.

12. 패턴 문서에 추가

notes/patterns/online_softmax.md에 아래를 추가하면 좋습니다.

# Pattern: Online Softmax Update

## CUDA

```cpp
float new_m = max(m, v);

float old_scale = expf(m - new_m);
float new_term  = expf(v - new_m);

s = s * old_scale + new_term;
m = new_m;

Observed SASS role pattern

// new_m = max(m, v)
FSETP.GT.FTZ ...
FSEL new_m, m, v, predicate

// old_scale = exp(m - new_m)
FADD old_delta, m, -new_m
FMUL old_delta, old_delta, log2(e)
MUFU.EX2 old_scale, old_delta

// new_term = exp(v - new_m)
FADD new_delta, v, -new_m
FMUL new_delta, new_delta, log2(e)
MUFU.EX2 new_term, new_delta

// s = s * old_scale + new_term
FFMA s, old_scale, s, new_term

Meaning

s = s * expf(old_m - new_m) + expf(v - new_m);
m = new_m;

Observed in

  • kernels/05_nn_ops/online_softmax_f32.cu

Materialization

  • m is held in register.
  • s is held in register.
  • old_scale and new_term are temporary register values.
  • No shared memory materialization is used in the update loop.

Key observation

Online softmax converts softmax from a shared-memory reduction/materialization pattern into a streaming register-accumulator update pattern.

---

# 13. 핵심 결론

이번 실험의 결론은 아주 큽니다.

```text
softmax_small_f32는 materialized softmax baseline이다.
online_softmax_f32는 register-accumulator softmax baseline이다.

더 직접적으로 말하면:

softmax_small:
    exp 전체를 shared memory에 저장하고 sum reduction을 수행한다.

online_softmax:
    exp 전체를 저장하지 않고,
    running max와 running sum을 update한다.

그리고 SASS에서 그 차이는 이렇게 보입니다.

softmax_small:
    STS / LDS / BAR.SYNC가 핵심

online_softmax:
    FSETP / FSEL / MUFU.EX2 / FFMA가 핵심

최종적으로 네가 세운 연구 프레임에 정확히 들어맞습니다.

CUDA source expression
→ compiler graph rewrite
→ SASS primitive
→ intermediate materialization 여부

이 실험에서 확인한 건:

FlashAttention식 online update의 축소판은
SASS 레벨에서 shared memory materialization이 아니라
register accumulator + FFMA update 패턴으로 읽힌다.

다음 단계는 둘 중 하나가 좋습니다.

1. online_softmax_f32 문서를 정리하고 1차 연구 노트 완성
2. flashattention_toy_f32 커널을 만들어서 QK score + online softmax + value accumulation까지 연결

내 추천은 1번입니다. 지금까지 결과를 먼저 문서화해야 다음 실험이 흔들리지 않습니다.