본문 바로가기

SASS_Probe

softmax_small_f32 분석

좋습니다. softmax_small_f32는 지금까지의 모든 패턴이 한 번에 섞여 있습니다.

결론부터 말하면:

softmax_small_f32는 중간값을 강하게 shared memory에 materialize하는 baseline softmax다.

구조는 정확히 이렇게 보입니다.

1. x[tid] load 또는 -FLT_MAX 초기화
2. shared memory에 저장
3. shared memory 기반 max reduction
4. max_v = sdata[0]
5. exp(x[tid] - max_v) 계산
6. exp 결과를 shared memory에 저장
7. shared memory 기반 sum reduction
8. sum_v = sdata[0]
9. reciprocal(sum_v) 계산
10. y[tid] = exp_value * reciprocal(sum_v)

특히 expf는 함수 호출이 아니라 아래 패턴으로 내려갔습니다.

FADD.FTZ R0, R3, -R0 ;
FMUL.FTZ R0, R0, 1.4426950216293334961 ;
MUFU.EX2 R8, R0 ;

즉:

expf(x - max_v)

가 내부적으로:

exp2((x - max_v) * log2(e))

형태로 바뀐 것입니다. x - max, * log2(e), MUFU.EX2 흐름이 실제 SASS에 보입니다.


softmax_small_f32 분석

문서 위치 추천:

notes/05_nn_ops/softmax_small_f32.md

1. CUDA 코드 구조

대상 커널은 대략 이런 형태입니다.

__global__ void softmax_small_f32_kernel(const float* __restrict__ x,
                                         float* __restrict__ y,
                                         int n) {
    extern __shared__ float sdata[];

    int tid = threadIdx.x;

    float v = -FLT_MAX;

    if (tid < n) {
        v = x[tid];
    }

    sdata[tid] = v;
    __syncthreads();

    // max reduction
    for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
        if (tid < stride) {
            float other = sdata[tid + stride];
            float self = sdata[tid];
            sdata[tid] = self > other ? self : other;
        }

        __syncthreads();
    }

    float max_v = sdata[0];
    __syncthreads();

    float e = 0.0f;

    if (tid < n) {
        e = expf(x[tid] - max_v);
        sdata[tid] = e;
    } else {
        sdata[tid] = 0.0f;
    }

    __syncthreads();

    // sum reduction
    for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
        if (tid < stride) {
            sdata[tid] += sdata[tid + stride];
        }

        __syncthreads();
    }

    float sum_v = sdata[0];
    __syncthreads();

    if (tid < n) {
        y[tid] = e / sum_v;
    }
}

high-level softmax:

y[i] = exp(x[i] - max(x)) / sum_j exp(x[j] - max(x));

2. SASS 전체 role map

이번 SASS는 크게 5개 phase로 나뉩니다.

Phase 1. load x or -FLT_MAX → sdata[tid]
Phase 2. shared memory max reduction
Phase 3. exp(x - max_v) 계산 → sdata[tid]
Phase 4. shared memory sum reduction
Phase 5. normalize and store

이 다섯 단계가 SASS에서 그대로 보입니다.


Phase 1. 입력 load 및 초기 shared memory 저장

핵심 구간:

/*0010*/ S2R R11, SR_TID.X ;
/*0030*/ MOV R0, 0xff7fffff ;
/*0050*/ ISETP.GE.AND P0, PT, R11.reuse, c[0x0][0x170], PT ;
/*0060*/ IMAD.WIDE R2, R11, R2, c[0x0][0x160] ;
/*0070*/ @!P0 LDG.E.CONSTANT R0, [R2.64] ;
/*00e0*/ STS [R11.X4], R0 ;
/*00f0*/ BAR.SYNC.DEFER_BLOCKING 0x0 ;

의미:

int tid = threadIdx.x;

float v = -FLT_MAX;

if (tid < n) {
    v = x[tid];
}

sdata[tid] = v;
__syncthreads();

여기서 MOV R0, 0xff7fffff는 -FLT_MAX에 해당하는 초기값으로 볼 수 있습니다.

R0 = -FLT_MAX
if (tid < n) R0 = x[tid]
sdata[tid] = R0

즉 max reduction을 위해 범위 밖 thread는 매우 작은 값을 넣습니다.


Phase 2. max reduction

핵심 구간:

/*0120*/ ISETP.GE.AND P3, PT, R11, R7, PT ;
/*0130*/ @!P3 LEA R0, R7, R4, 0x2 ;
/*0140*/ @!P3 LDS R6, [R11.X4] ;
/*0160*/ @!P3 LDS R9, [R0] ;
/*0170*/ @!P3 FSETP.GT.FTZ.AND P2, PT, R6, R9, PT ;
/*0180*/ @!P3 FSEL R6, R6, R9, P2 ;
/*01a0*/ @!P3 STS [R11.X4], R6 ;
/*01b0*/ BAR.SYNC.DEFER_BLOCKING 0x0 ;
/*01c0*/ @P2 BRA 0x120 ;

의미:

if (tid < stride) {
    float self = sdata[tid];
    float other = sdata[tid + stride];
    sdata[tid] = self > other ? self : other;
}

__syncthreads();

여기서 max는 FMNMX가 아니라:

FSETP + FSEL

로 구현되었습니다.

FSETP.GT.FTZ.AND P2, PT, R6, R9, PT ;
FSEL R6, R6, R9, P2 ;

즉:

R6 = (R6 > R9) ? R6 : R9;

입니다. shared memory에서 self, other를 읽고 비교/선택 후 다시 sdata[tid]에 저장하는 tree reduction입니다.

이 단계의 materialization은 명확합니다.

LDS self
LDS other
FSETP/FSEL max
STS sdata[tid]
BAR.SYNC

즉 max partial result가 매 stride마다 shared memory에 저장됩니다.


Phase 3. exp(x - max_v)

max reduction이 끝나면:

/*01d0*/ LDS R0, [RZ] ;
/*01f0*/ BAR.SYNC.DEFER_BLOCKING 0x0 ;
/*0200*/ @P0 BRA 0x250 ;
/*0210*/ LDG.E.CONSTANT R3, [R2.64] ;
/*0220*/ FADD.FTZ R0, R3, -R0 ;
/*0230*/ FMUL.FTZ R0, R0, 1.4426950216293334961 ;
/*0240*/ MUFU.EX2 R8, R0 ;
/*0250*/ BSYNC B0 ;
/*0260*/ STS [R11.X4], R8 ;
/*0270*/ BAR.SYNC.DEFER_BLOCKING 0x0 ;

의미:

float max_v = sdata[0];
__syncthreads();

float e = 0.0f;

if (tid < n) {
    e = expf(x[tid] - max_v);
}

sdata[tid] = e;
__syncthreads();

여기서 핵심은 expf가 어떻게 내려갔는지입니다.

FADD.FTZ R0, R3, -R0 ;
FMUL.FTZ R0, R0, 1.4426950216293334961 ;
MUFU.EX2 R8, R0 ;

이를 high-level로 바꾸면:

R0 = x[tid] - max_v;
R0 = R0 * log2(e);
R8 = exp2(R0);

즉:

R8 = expf(x[tid] - max_v);

입니다. 1.4426950216293334961은 log2(e)입니다. MUFU.EX2는 base-2 exponential special function unit 계열 명령으로 볼 수 있습니다.

중요한 materialization:

STS [R11.X4], R8 ;

즉 e = expf(x - max)가 shared memory에 저장됩니다.

e는 register에만 남지 않는다.
sdata[tid]에 materialize된다.

Phase 4. sum reduction

핵심 구간:

/*0290*/ ISETP.GE.AND P1, PT, R11, R5, PT ;
/*02a0*/ @!P1 IMAD R0, R5, 0x4, R4 ;
/*02b0*/ @!P1 LDS R2, [R11.X4] ;
/*02c0*/ SHF.R.U32.HI R5, RZ, 0x1, R5 ;
/*02d0*/ @!P1 LDS R3, [R0] ;
/*02e0*/ @!P1 FADD.FTZ R2, R2, R3 ;
/*02f0*/ @!P1 STS [R11.X4], R2 ;
/*0300*/ BAR.SYNC.DEFER_BLOCKING 0x0 ;
/*0310*/ ISETP.NE.AND P1, PT, R5, RZ, PT ;
/*0320*/ @P1 BRA 0x290 ;

의미:

for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
    if (tid < stride) {
        sdata[tid] += sdata[tid + stride];
    }

    __syncthreads();
}

이건 앞에서 본 reduce_sum_f32와 거의 같은 공유 메모리 tree reduction입니다.

핵심 패턴:

LDS
LDS
FADD
STS
BAR.SYNC
BRA

즉 sum_v = sum(exp(...))도 register-only가 아니라 shared memory에 계속 materialize됩니다.


Phase 5. normalize and store

마지막 구간:

/*0330*/ LDS R0, [RZ] ;
/*0340*/ BAR.SYNC.DEFER_BLOCKING 0x0 ;
/*0350*/ @P0 EXIT ;
/*0360*/ MUFU.RCP R5, R0 ;
/*0380*/ IMAD.WIDE R2, R11, R2, c[0x0][0x168] ;
/*0390*/ FMUL.FTZ R5, R5, R8 ;
/*03a0*/ STG.E [R2.64], R5 ;
/*03b0*/ EXIT ;

의미:

float sum_v = sdata[0];
__syncthreads();

if (tid >= n) {
    return;
}

float inv_sum = 1.0f / sum_v;
y[tid] = e * inv_sum;

division이 직접 나오는 게 아니라:

MUFU.RCP R5, R0 ;
FMUL.FTZ R5, R5, R8 ;

형태로 나왔습니다.

즉:

R5 = reciprocal(sum_v);
R5 = R5 * e;

입니다.

CUDA source의:

y[tid] = e / sum_v;

가 SASS에서는:

y[tid] = e * rcp(sum_v);

로 바뀐 것입니다.


3. intermediate materialization 관점

이 커널은 materialization 관점에서 아주 좋은 baseline입니다.

materialized 되는 것

1. max reduction 입력

sdata[tid] = v;

SASS:

STS [R11.X4], R0 ;

2. max reduction partial result

sdata[tid] = max(sdata[tid], sdata[tid + stride]);

SASS:

LDS
LDS
FSETP
FSEL
STS

3. exp result

sdata[tid] = e;

SASS:

MUFU.EX2 R8, R0 ;
STS [R11.X4], R8 ;

4. sum reduction partial result

sdata[tid] += sdata[tid + stride];

SASS:

LDS
LDS
FADD
STS

따라서 이 softmax는:

x
→ shared memory for max reduction
→ max_v
→ exp(x - max_v)
→ shared memory for sum reduction
→ sum_v
→ output

구조입니다.

즉 중간값이 여러 번 shared memory에 저장됩니다.


4. register 유지되는 값

반대로 register에 유지되는 값도 있습니다.

R8 = e = exp(x[tid] - max_v)

흥미롭게도 R8은 STS [R11.X4], R8로 shared memory에 저장된 뒤에도 마지막 normalize에서 다시 사용됩니다.

마지막에:

FMUL.FTZ R5, R5, R8 ;

가 나오기 때문입니다.

즉 이 커널은:

e를 shared memory에 저장한다.
동시에 현재 thread의 e는 R8 register에도 남아 있다.

이렇게 볼 수 있습니다.

하지만 sum reduction은 shared memory에 저장된 e들을 읽어야 하므로, operator 전체 관점에서는 e가 materialized된 것이 맞습니다.


5. reduce_sum_f32와 비교

reduce_sum_f32:

LDG
STS
BAR
loop:
    LDS
    LDS
    FADD
    STS
    BAR
STG

softmax_small_f32:

LDG or -FLT_MAX
STS
BAR

max loop:
    LDS
    LDS
    FSETP
    FSEL
    STS
    BAR

LDS max_v

exp:
    LDG
    FADD x-max
    FMUL *log2(e)
    MUFU.EX2
    STS e
    BAR

sum loop:
    LDS
    LDS
    FADD
    STS
    BAR

LDS sum_v
MUFU.RCP
FMUL e*rcp
STG

즉 softmax_small은 reduce_sum 패턴을 두 번 포함합니다.

max reduction 1회
sum reduction 1회

그리고 그 사이에:

exp transform

이 끼어 있습니다.


6. 패턴 문서에 추가

notes/patterns/softmax_baseline.md에 아래처럼 정리하면 됩니다.

# Pattern: Baseline Shared-Memory Softmax

## CUDA

```cpp
max_v = reduce_max(x);
e = exp(x[i] - max_v);
sum_v = reduce_sum(e);
y[i] = e / sum_v;

Observed SASS role pattern

// initial load
LDG x
STS sdata[tid], x
BAR.SYNC

// max reduction
loop:
    LDS self
    LDS other
    FSETP.GT
    FSEL max
    STS sdata[tid], max
    BAR.SYNC
    BRA loop

// exp
LDS max_v, sdata[0]
LDG x
FADD x_minus_max, x, -max_v
FMUL scaled, x_minus_max, log2(e)
MUFU.EX2 e, scaled
STS sdata[tid], e
BAR.SYNC

// sum reduction
loop:
    LDS self
    LDS other
    FADD sum
    STS sdata[tid], sum
    BAR.SYNC
    BRA loop

// normalize
LDS sum_v, sdata[0]
MUFU.RCP inv_sum, sum_v
FMUL out, e, inv_sum
STG y, out

Materialization

  • max reduction partials are materialized in shared memory.
  • exp result e is materialized in shared memory.
  • sum reduction partials are materialized in shared memory.
  • final output is stored in global memory.

Key observation

This baseline softmax exposes the full shared-memory materialization path:

x → shared max buffer → exp buffer → shared sum buffer → y

This is useful as a contrast target for online softmax / FlashAttention-style kernels, where the goal is to avoid materializing the full intermediate probability matrix.

---

# 7. 핵심 결론

`softmax_small_f32`는 우리가 세운 분석 프레임에 아주 잘 맞습니다.

```text
CUDA source expression:
    softmax(x) = exp(x - max(x)) / sum(exp(x - max(x)))

compiler graph rewrite:
    shared-memory max reduction
    exp via exp2 transform
    shared-memory sum reduction
    reciprocal + multiply normalization

SASS primitive:
    LDS / STS / BAR.SYNC
    FSETP + FSEL
    FADD / FMUL
    MUFU.EX2
    MUFU.RCP
    STG

intermediate materialization:
    매우 많음.
    max partials, exp values, sum partials가 shared memory에 저장됨.

중요한 문장으로 정리하면 이겁니다.

softmax_small_f32는 softmax를 operator-level로 직접 보여주는 좋은 baseline이다.
SASS 상에서 max reduction, exp transform, sum reduction, normalization이 명확히 분리되어 보이며,
중간 exp 값은 shared memory에 materialize된다.

다음은 online_softmax_f32를 보면 됩니다.
거기서는 지금과 비교해서 핵심 질문이 하나로 좁혀집니다.

online update에서는 exp 전체 벡터를 shared memory에 저장하지 않고,
running max / running sum을 register accumulator로 유지하는가?

'SASS_Probe' 카테고리의 다른 글

1차 연구 노트 : CUDA Source에서 SASS Primitive 로 보는 Operator Materialization  (0) 2026.06.16
online_softmax_f32 분석  (0) 2026.06.14
reduce_sum_f32 분석  (0) 2026.06.14
clamp_f32 분석  (0) 2026.06.14
relu_f32 분석  (0) 2026.06.14