FNet: 푸리에 변환으로 어텐션을 대체한 효율적 토큰 믹싱 모델
Google Research · 2021-05-09 · SSM · Apache-2.0
개요
FNet은 2021년 Google Research가 발표한 모델로, Transformer의 self-attention을 푸리에 변환(Fourier Transform)으로 완전히 대체한 실험적 아키텍처이다. 이 모델의 핵심 주장은 명확하다. 시퀀스 내 토큰 간 정보 혼합(token mixing)에 반드시 학습 가능한 어텐션 메커니즘이 필요한 것은 아니라는 것이다. 2D FFT(Fast Fourier Transform)를 적용하면 학습 파라미터 없이도 BERT 성능의 92~97%를 유지할 수 있음을 실험적으로 입증했다.
FFT는 순수 수학적 연산으로 학습 가능한 파라미터가 전혀 없다. 이로 인해 GPU/TPU에서 극도로 빠르게 동작하며, BERT 대비 훈련 속도는 최대 7배, 추론 속도는 2~4배 빠르다. FNet은 attention-free 모델링의 가능성을 탐색한 선구적 연구로, 이후 등장한 S4, Mamba 등 SSM(State Space Model) 계열 연구에 “어텐션이 유일한 해법은 아니다”라는 중요한 메시지를 전달했다.
Transformer의 self-attention은 복잡도를 가지므로 시퀀스 길이가 길어질수록 계산 비용이 급격히 증가한다. FNet은 이 병목을 FFT의 복잡도로 해소하면서도, 대부분의 NLU(Natural Language Understanding) 태스크에서 실용적 수준의 성능을 유지할 수 있음을 보여준 최초의 대규모 실험이다.
아키텍처 상세
FNet의 아키텍처는 Transformer와 거의 동일한 구조를 유지하되, Multi-Head Attention 레이어를 FFT 레이어로 교체한 것이 핵심이다. 각 FNet 블록은 다음 두 단계로 구성된다.
FFT 토큰 믹싱 레이어
입력 행렬 에 대해 2D DFT(Discrete Fourier Transform)를 적용한다. 먼저 시퀀스 차원()에 대해 FFT를 수행하여 토큰 간 전역적 정보 혼합을 달성하고, 이어서 히든 차원()에 대해 FFT를 수행하여 특성 간 상호작용을 모델링한다.
여기서 는 복소수 출력의 실수부만 추출하는 연산이다. DFT 자체의 정의를 전개하면 시퀀스 차원에 대해 다음과 같다.
이 변환은 시간 영역의 시퀀스를 주파수 영역으로 변환하며, 모든 위치의 정보를 전역적으로 혼합한다. FFT 알고리즘은 이를 에 계산하므로 어텐션의 보다 효율적이다.
피드포워드 네트워크
FFT 출력에 기존 Transformer와 동일한 2층 FFN(Feed-Forward Network)을 적용한다. GELU 활성화 함수와 LayerNorm을 사용하며, 잔차 연결(residual connection)도 동일하게 유지된다.
중요한 점은 FFT 레이어에 학습 가능한 파라미터가 전혀 없다는 것이다. 모든 학습은 FFN과 임베딩 레이어에서만 이루어진다. 위치 정보는 learnable absolute positional embedding으로 별도 처리한다.
SSM 관점에서의 해석
FNet의 FFT 연산은 SSM(State Space Model)의 컨볼루션 모드와 깊은 연관이 있다. S4에서 상태 공간 모델의 이산화된 커널을 시퀀스에 적용할 때 FFT를 활용하는데, FNet은 이 커널 자체를 학습하지 않고 단순 DFT 기저를 그대로 사용하는 것으로 해석할 수 있다. 연속 시간 SSM의 출력은 다음과 같이 컨볼루션으로 표현된다.
FNet은 이 커널 를 학습 없이 DFT 기저 함수로 고정한 특수 케이스로 볼 수 있다.
핵심 혁신
FNet의 핵심 혁신은 세 가지로 요약된다.
첫째, 파라미터 없는 토큰 믹싱이다. 기존 어텐션은 Q, K, V 프로젝션 행렬로 토큰 간 상호작용을 학습하지만, FNet은 수학적 변환만으로 이를 달성한다. 이는 모델 파라미터 수를 크게 줄이면서도 충분한 성능을 유지할 수 있음을 보여준다.
둘째, 전역적 정보 혼합이다. FFT는 본질적으로 모든 주파수 성분을 동시에 처리하므로, 시퀀스 내 모든 위치의 정보가 단일 연산으로 혼합된다. 이는 local attention이 아닌 global mixing을 제공한다.
셋째, 하드웨어 최적화 가능성이다. FFT는 GPU/TPU에서 고도로 최적화된 라이브러리(cuFFT, JAX FFT 등)가 이미 존재하므로, 별도의 커스텀 커널 개발 없이도 높은 처리량을 달성할 수 있다.
벤치마크/성능
| 벤치마크 | FNet | BERT-Base | 상대 성능 |
|---|---|---|---|
| GLUE 평균 | ~80.5 | ~82.2 | 97.9% |
| MNLI | 76.7 | 84.6 | 90.7% |
| QQP | 88.4 | 91.1 | 97.0% |
| SST-2 | 92.2 | 93.5 | 98.6% |
| 훈련 속도(TPU) | 7x faster | 1x | - |
| 추론 속도(GPU) | 2~4x faster | 1x | - |
GLUE 벤치마크 전체 평균 기준으로 BERT-Base의 약 97~98% 성능을 유지하면서도 훈련 비용은 최대 7배 절감된다. 특히 SST-2 같은 단일 문장 분류 태스크에서는 거의 동등한 성능을 보인다.
| 모델 | 토큰 믹싱 | 복잡도 | 학습 파라미터 | 특징 |
|---|---|---|---|---|
| FNet | 2D FFT | 없음(FFT) | 가장 빠른 학습 | |
| Transformer | Self-Attention | Q,K,V 프로젝션 | 최고 성능 | |
| S4 | SSM 커널 | A,B,C 행렬 | 장거리 의존성 | |
| Hyena | 암묵적 컨볼루션 | FFN 생성 필터 | Attention-free |
학습
FNet의 학습은 BERT 사전학습 절차를 최대한 준수하여 공정한 비교를 보장했다. C4 데이터셋으로 사전학습하며, SentencePiece 토크나이저(32K vocab)를 사용한다. TPUv3-8 환경에서 배치 크기 256, 최대 시퀀스 길이 512 토큰으로 학습했다. Adam 옵티마이저를 사용하며 웜업 스텝 10,000을 적용했다. MLM(Masked Language Modeling) 목표로 학습되며, 파인튜닝은 GLUE 태스크에서 수행했다.
다음은 FNet의 FFT 토큰 믹싱 레이어를 PyTorch로 구현한 예시이다.
import torch
import torch.nn as nn
import torch.fft
class FNetBlock(nn.Module):
def __init__(self, d_model, d_ff, dropout=0.1):
super().__init__()
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(d_ff, d_model),
nn.Dropout(dropout),
)
def forward(self, x):
# FFT 토큰 믹싱 (학습 파라미터 없음)
fft_out = torch.fft.fft2(x).real
x = self.norm1(x + fft_out) # 잔차 연결
# 피드포워드 네트워크
x = self.norm2(x + self.ffn(x))
return x이 구현에서 볼 수 있듯이 torch.fft.fft2는 시퀀스 차원과 히든 차원에 대해 동시에 FFT를 수행하며, .real로 실수부만 추출한다. 전체 FFT 레이어에 학습 가능한 파라미터가 전혀 없다는 점이 핵심이다.
관련 모델
FNet은 Transformer 대체 모델 중 가장 단순한 접근법으로, 복잡한 SSM 수식이나 선택적 메커니즘 없이도 합리적인 성능을 달성한다는 점에서 의의가 있다. FFT가 입력 내용에 무관한 고정 변환이라는 한계(LTI 특성)는 이후 S4에서도 동일하게 관찰되었으며, 이는 Mamba의 선택적 메커니즘(Selective Mechanism)으로 해결되었다. FNet은 “어텐션만이 답이 아니다”라는 메시지로 SSM 연구 생태계의 성장에 기여한 중요한 이정표이다.
참고 자료
관련 문서
- Transformer — 영감