Selective State Space Model을 구현 관점에서 이해하기
요약
본 글은 Mamba 모델의 핵심 메커니즘인 Selective State Space Model(SSM)을 구현 관점에서 깊이 있게 분석합니다. SSM의 재귀적 특성과 컨볼루션적 이중성을 설명하며, 입력 의존적인 업데이트가 계산 그래프를 어떻게 변화시키는지 추적하는 것이 목표입니다.
핵심 포인트
- Mamba는 Selective SSM, Selective Scan, Mamba block 세 계층으로 구성됩니다.
- SSM은 재귀(recurrent)와 컨볼루션(convolution)의 이중성을 가집니다.
- Selective SSM은 입력에 따라 상태를 결정하는 핵심 메커니즘입니다.
- 글은 Mamba 원 논문과 공식 repository를 근거로 설명합니다.
Mamba의 구현 코드를 읽다 보면, selective_scan, d_state, dt_proj와 같은 이름들이 한 번에 등장합니다. 어떤 것이 모델의 표현력을 담당하고, 어떤 것이 계산 속도를 높이는 메커니즘인지 구분하지 않으면 전체 그림을 파악하기 어렵습니다.
먼저 결론부터 말씀드리자면, Mamba는 다음 세 가지 계층으로 나누어 이해할 수 있습니다.
- Selective SSM: 입력에 따라 무엇을 state로 남길지 결정합니다.
- Selective Scan: 입력 의존적인 재귀(recurrent) 업데이트를 시퀀스 방향으로 효율적으로 계산합니다.
- Mamba block: local Conv1D, Selective SSM, gate, residual을 통합합니다.
이 글에서는 고정된 SSM 수식에서 출발하여, 입력 의존화(input-dependent)가 계산 그래프를 어떻게 변화시키는지 추적해 나갑니다. 주요 근거는 Mamba 원 논문과 저자의 공식 repository입니다.
고정 SSM을 이산화하기
SSM (State Space Model: 상태 공간 모델)은 입력 $\mathbf{x}(t)$를 받아 다음과 같은 재귀(recurrent) 형태로 업데이트됩니다.
$$\mathbf{h}(t) = \mathbf{\Phi} \mathbf{h}(t-1) + \mathbf{B} \mathbf{x}(t)$$
$$\mathbf{y}(t) = \mathbf{C} \mathbf{h}(t) + \mathbf{D} \mathbf{x}(t)$$
실제 Mamba에서는 배치(batch)를 다음과 같이 처리합니다.
| tensor | 개념적인 shape | 의미 |
|---|---|---|
| input | 각 시점의 channel 표현 | |
| state | 채널별 재귀 상태 | |
| state의 기본적인 시간 발전 | ||
| token・channel별 스텝 폭 | ||
| 구현에 따라 broadcast 가능한 sequence 의존 shape | 쓰기/읽기 parameter |
repository의 kernel interface와 내부 레이아웃은 버전이나 구현에 따라 다를 수 있습니다. 여기서는 수식상의 축(axis)을 보여주는 것이며, 특정 API의 물리적 레이아웃을 그대로 나타내는 것은 아닙니다.
고정 SSM은 재귀이면서 컨볼루션이다

이러한 이중성(duality) 덕분에, 생성 시에는 이전 state만을 사용하여 한 스텝씩 업데이트하고, 학습 시에는 전체 시퀀스를 컨볼루션으로 병렬 계산할 수 있습니다. S4 논문은
다만, 고정된 커널(fixed kernel)로는 입력 내용에 따라 업데이트 규칙을 바꿀 수 없습니다. 언어와 같은 이산 데이터에서는
다음 코드는 스칼라(scalar) 상태에 한정하여 교육용으로 구현한 것입니다. 실제 Mamba 커널이 처리하는 배치(batch), 채널(channel), 상태 축(state axis), 병렬 스캔(parallel scan), 커널 융합(kernel fusion)은 생략했습니다.
from collections.abc import Sequence
from math import exp
def selective_scan_scalar(
...
레퍼런스 구현에서 다음 대응 관계를 확인할 수 있습니다.
a_bar:
이전 상태(state)를 얼마나 유지할지 -b_bar * x_t:
현재 입력(input)을 어떻게 기록할지 -c_t * state:
현재 시점에서 무엇을 읽어낼지
프로덕션 구현에서는 이 Python 루프를 그대로 사용하지 않습니다. 텐서(tensor)를 전개한 거대한 중간 상태(intermediate state)를 HBM에 저장하면 메모리 트래픽(memory traffic)이 증가하기 때문에, Mamba 논문은 커널 융합, 병렬 스캔, 재계산(recomputation)을 결합합니다.

재계산은 backward 과정에서 일부 중간값을 다시 계산하여 HBM으로의 물질화(materialization)를 피하는 트레이드오프입니다. 산술량뿐만 아니라 HBM과 온칩 SRAM 사이의 데이터 이동을 설계 대상으로 삼습니다.
Mamba 블록은 SSM만이 아니다
원래 Mamba 블록에서는 입력을 메인 브랜치(main branch)와 게이트 브랜치(gate branch)로 나눕니다.

메인 브랜치의 개념적인 순서는 다음과 같습니다.
게이트 브랜치는 별도의 선형 투영(linear projection)과 SiLU를 거쳐 메인 브랜치의 출력에 요소별 곱(element-wise product)으로 곱해집니다. 그 후에 출력 투영(output projection)과 잔차 연결(residual connection)이 있습니다.
| 컴포넌트 | 주요 역할 | 구현을 읽을 때 확인할 점 |
|---|---|---|
| causal Conv1D | 인접 토큰의 로컬 믹싱 | 커널 폭(kernel width)과 증분 버퍼(incremental buffer) |
| Selective SSM | 긴 방향으로의 상태 전파 | |
| gate | 출력 채널의 내용에 따른 종속 제어 | 브랜치의 투영(projection)과 활성화 함수(activation) |
| residual | 블록 간 정보 경로 | 정규화(normalization) 위치 |
Mamba는 Attention 레이어를 Selective SSM으로 일대일로 대체한 구조가 아닙니다. 로컬 컨볼루션과 게이트를 포함한 블록 전체에서 하나의 시퀀스 믹서(sequence mixer)를 구성합니다.
training/prefill과 decode의 상태 관리

training 및 prefill에서는 시퀀스 전체가 알려져 있으므로, 여러 위치에 대한 부분 변환을 병렬 스캔으로 합성합니다. 총 연산량은 시퀀스 길이 $\text{L}$입니다.
자기회귀(autoregressive) 디코드는 토큰당 다음 두 가지를 업데이트합니다.
- 각 레이어의 SSM 상태
- causal Conv1D의 짧은 상태 버퍼
과거 토큰별 Key/Value를 유지하는 Transformer의 KV 캐시와 달리, 이러한 상태량은 시퀀스 길이 $\text{L}$에 비례합니다.
Transformer와의 트레이드오프는 직접 접근인가 압축인가
| 관점 | Dense Attention | Mamba |
|---|---|---|
| 시퀀스 전체의 계산 | $O(\text{L}^2)$ | $O(\text{L} \cdot d^2)$ |
| ... | ||
| 고정 길이 상태는 스트리밍(streaming) 및 긴 시퀀스에서 유리하지만, 과거의 모든 것을 손실 없이 유지할 수는 없습니다. Attention은 토큰별 표현을 남기는 대신, 계산량과 메모리가 길이에 비례하여 증가합니다. |
따라서, 이론적 복잡도만으로 아키텍처를 선택하는 것은 불충분합니다. 정확한 복사(copy)나 과거 위치의 검색(retrieval)이 중요하다면 상태 병목(state bottleneck)을 측정하고, 긴 신호의 순차 처리가 중심이라면 고정 크기 상태의 이점을 측정해야 합니다. Attention과 SSM을 결합하는 하이브리드 방식도 이러한 트레이드오프에 대한 자연스러운 선택지입니다.
구현을 읽는 체크리스트
공식 구현이나 파생 모델을 읽을 때는 다음 순서로 따라가면 헷갈리지 않습니다.
구현을 읽는 체크리스트
공식 구현이나 파생 모델을 읽을 때는 다음 순서로 따라가면 헷갈리지 않습니다.
- input/output tensor의 axis 순서와 channel 전개율(channel expansion rate)을 확인한다
- A에서 parameterization과 안정성의 제약 조건을 확인한다
- $\Delta$에 양수 제약(positive constraint)을 주는 변환을 확인한다
- B, C가 어떤 axis로 입력 의존적(input dependent)으로 broadcast되는지 확인한다
- discretization이 kernel 내 어디에서 이루어지는지 확인한다
- full-sequence scan과 incremental update의 진입점(entry point)을 분리한다
- Conv1D state와 SSM state의 초기화 및 업데이트를 확인한다
- fallback 구현과 optimized kernel로 수치적 차이를 테스트한다
- 시퀀스 길이(sequence length)를 스윕하여 처리량(throughput)과 최대 메모리(peak memory)를 개별적으로 측정한다
- copy, retrieval, streaming 등 태스크별로 품질(quality)을 측정한다
원 논문의 최대 5배에 달하는 추론 처리량(inference throughput)은 논문 내에서 보고된 하드웨어, 모델, 배치, 시퀀스 조건에서의 수치입니다. 현재의 임의 트랜스포머 구현에 그대로 적용하지 말고, 자신의 shape과 서빙 조건에서 프로파일링해야 합니다.
요약 (Summary)
Mamba의 selection mechanism은,
구현 관점에서는 수식상의 Selective SSM, 계산 알고리즘으로서의 Selective Scan, Conv1D와 게이트(gate)를 포함하는 Mamba 블록을 분리하여 추적하는 것이 핵심입니다. 또한, full-sequence 계산과 incremental decode에서는 사용되는 경로와 유지되는 state가 다릅니다.
선형 시간(linear time)과 고정 길이 state는 계산상의 장점입니다. 반면, 과거를 유한한 state로 압축하는 정보 병목 현상(information bottleneck)은 여전히 존재합니다. 트랜스포머와의 비교에서는 속도의 우위가 아니라, 직접 접근(direct access)과 압축된 state 간의 상충 관계(trade-off)를 태스크 위에서 측정해야 합니다.
SSM의 연속 시간식, 이산화의 구체적인 계산, 용어의 계층 구조, 트랜스포머와의 비교를 처음부터 확인하고 싶은 분은 개인 블로그에 올라온 Mamba / SSM 완전판도 참고해 주세요.
참고 자료
논의 (Discussion)

AI 자동 생성 콘텐츠
본 콘텐츠는 Zenn ML의 원문을 AI가 자동으로 요약·번역·분석한 것입니다. 원 저작권은 원저작자에게 있으며, 정확한 내용은 반드시 원문을 확인해 주세요.
원문 바로가기