Chapter 19

FFT로 합성곱하기

주파수 영역에서 빠르게 곱하기
핵심 질문

합성곱이 왜 FFT 뒤에는 자리별 곱셈이 될까?

FFT로 합성곱하기

핵심 질문

FFT는 실제로 다항식 곱셈을 어떻게 빠르게 만들까?

교과서 설명

FFT를 이용한 합성곱은 세 단계로 볼 수 있다.

주파수 영역에서는 합성곱이 자리별 곱셈으로 바뀐다.

왜 곱셈이 되는지 핵심만 보면 다음과 같다. 먼저 두 배열 뒤에 을 충분히 붙여 같은 길이 로 맞추었다고 보자. 이 길이는 선형 합성곱 결과가 말려 들어가지 않도록 을 만족해야 한다.

이제 이고, 합성곱은 라고 두자.

가 되는 모든 곱 를 모은 값이다. 그래서 모든 를 돌며 더하는 것은 결국 모든 쌍 를 한 번씩 더하는 것과 같다. 이 생각으로 로 펼치면 DFT 합은 이렇게 정리된다.

시간 영역에서는 여러 위치를 밀어 더해야 하지만, 주파수 영역에서는 같은 끼리만 곱하면 된다.

마지막으로 역 FFT를 하면 원래 영역의 합성곱을 얻는다. 앞 장에서 본 것처럼 IFFT는 회전 방향을 반대로 쓰고, 마지막에 길이 로 나누는 역변환이다.

여기서 IFFT는 주파수 영역의 값 를 다시 계수 배열 로 되돌리는 단계다. FFT와 자리별 곱셈까지만 하면 아직 답은 주파수 영역에 남아 있다.

이렇게 뒤에 을 충분히 붙이는 과정을 zero padding이라고 한다.

두 배열 길이가 , 이면 결과 길이가 이므로 FFT 길이는 적어도 이만큼 필요하다.

보통 FFT가 편하게 동작하도록 의 거듭제곱으로 잡는다.

이때 전체 계산량은 길이 짜리 FFT 두 번, 역 FFT 한 번, 그리고 자리별 곱셈 번으로 볼 수 있다.

직접 합성곱의 또는 길이가 비슷할 때의 보다 큰 입력에서 훨씬 유리해진다.

또 한 가지 주의할 점이 있다. DFT 세계에서 그냥 곱하면 원형으로 말려 들어가는 순환 합성곱이 된다. zero padding을 충분히 해야 우리가 원하는 선형 합성곱과 같아진다.

즉 이 장의 zero padding은 합성곱 결과를 안전하게 담기 위한 장치다. 샘플링된 신호 뒤에 을 붙인다고 원래 없던 정보가 생기는 것은 아니다.

합성곱 정리의 조건

DFT 길이를 로 고정하면 모든 인덱스는 로 나눈 나머지처럼 작동한다. 그래서 DFT 영역에서 곱한 뒤 IFFT하면 기본적으로 길이 의 순환 합성곱이 나온다.

우리가 다항식 곱셈에서 원하는 것은 선형 합성곱이다.

두 결과가 같아지려면 선형 합성곱의 모든 결과 칸이 부터 안에 들어와야 한다. 그래서

이 필요하다. 이 조건을 만족하도록 뒤에 을 붙이면, 원래는 앞으로 말려 들어갈 항들이 안전한 뒤쪽 칸에 남는다. 반대로 이 작으면 뒤쪽 계수가 앞쪽 계수에 더해져 버리므로, 겉으로는 길이 짜리 답이 나오더라도 다항식 곱셈 결과는 틀린다.

다항식 관점에서는 같은 사실을 이렇게 말할 수 있다. 충분히 많은 점에서 를 평가하고, 각 점에서 값을 곱해 의 값을 얻은 뒤, 역변환으로 의 계수를 복원한다. FFT는 평가와 복원을 빠르게 해 주고, 자리별 곱셈은 값 표현에서 다항식 곱셈을 수행하는 단계다.

예를 들어 길이 로만 계산하면 결과가 두 칸 안에서 말려 들어간다.

하지만 길이 순환 합성곱으로 보면 마지막 이 앞 칸으로 돌아와 다음처럼 섞인다.

여기서 은 사라진 것이 아니라 번 칸으로 접혀 들어가 과 더해진다. 이 현상이 aliasing이다. 그래서 길이 이상의 공간을 만들기 위해 을 붙여야 한다.

직관 비유

복잡한 퍼즐을 직접 맞추기 어렵다면, 퍼즐 조각을 색깔별로 분류한 뒤 같은 색끼리 처리하고 다시 합치는 방법을 쓸 수 있다. FFT는 문제를 “계산하기 쉬운 세계”로 잠깐 옮기는 도구다.

예제

두 배열 , 의 합성곱은 다음과 같다.

FFT로 계산할 때는 길이 이상이 필요하므로 길이 로 맞출 수 있다.

계산 순서는 다음과 같다.

컴퓨터에서는 복소수 소수 오차가 생길 수 있으므로 마지막 결과가 처럼 나오면 가까운 정수로 반올림한다.

손풀이 체크

  1. 길이 배열 두 개의 선형 합성곱 결과 길이는?

    답 보기

    2+21=32+2-1=3

  2. 을 길이 순환 합성곱으로 보면 무엇이 앞 칸으로 말려 들어갈까?

    답 보기

    마지막 값 88

  3. FFT 합성곱의 세 핵심 단계는?

    답 보기

    FFT, 자리별 곱셈, IFFT

  4. 를 길이 로 zero padding해서 계산하면 IFFT 결과의 앞 세 칸은 무엇이어야 할까?

    답 보기

    선형 합성곱 결과인 [3,10,8][3,10,8]이어야 한다. 네 번째 칸은 padding 때문에 생긴 여분 칸이다.

마무리

DFT는 삼각함수, 복소수, 벡터, 좌표평면이 “회전으로 신호를 읽는 법”으로 이어진 결과다. FFT는 그 DFT를 빠르게 계산해서 합성곱과 다항식 곱셈에 사용할 수 있게 해 주는 알고리즘이다.

이번 장에서 기억할 3문장

  1. FFT 합성곱은 FFT, 자리별 곱셈, IFFT 순서로 선형 합성곱을 계산한다.
  2. DFT 길이 L이 부족하면 결과가 순환 합성곱처럼 앞쪽으로 말려 들어간다.
  3. zero padding은 결과를 담을 공간을 만드는 장치이지 새 정보를 만드는 장치가 아니다.

C++ Practice

C++로 확인하기

FFT로 FFT(a)FFT(b)\text{FFT}(a)\cdot\text{FFT}(b)를 계산해 다항식 곱셈 결과 얻기

#include <cmath>
#include <complex>
#include <iostream>
#include <vector>

using namespace std;

using Complex = complex<double>;

void fft(vector<Complex>& a, bool inverse) {
    int N = static_cast<int>(a.size());
    if (N == 1) {
        return;
    }

    vector<Complex> even(N / 2), odd(N / 2);
    for (int i = 0; i < N / 2; ++i) {
        even[i] = a[2 * i];
        odd[i] = a[2 * i + 1];
    }

    fft(even, inverse);
    fft(odd, inverse);

    const double pi = acos(-1.0);
    double sign = inverse ? 1.0 : -1.0;
    Complex omega = polar(1.0, sign * 2.0 * pi / N);
    Complex power = 1;

    for (int k = 0; k < N / 2; ++k) {
        Complex t = power * odd[k];
        a[k] = even[k] + t;
        a[k + N / 2] = even[k] - t;
        power *= omega;
    }
}

vector<int> multiply(vector<int> a, vector<int> b) {
    int needed = static_cast<int>(a.size() + b.size() - 1);
    int N = 1;
    while (N < needed) {
        N *= 2;
    }

    vector<Complex> A(N), B(N);
    for (int i = 0; i < static_cast<int>(a.size()); ++i) A[i] = a[i];
    for (int i = 0; i < static_cast<int>(b.size()); ++i) B[i] = b[i];

    fft(A, false);
    fft(B, false);
    for (int i = 0; i < N; ++i) {
        A[i] *= B[i];
    }
    fft(A, true);

    vector<int> result(needed);
    for (int i = 0; i < needed; ++i) {
        result[i] = static_cast<int>(round(A[i].real() / N));
    }
    return result;
}

int main() {
    auto result = multiply({1, 2}, {3, 4});
    for (int value : result) {
        cout << value << ' ';
    }
    cout << "\n";
}

연습: a=[1,2,3]a=[1,2,3], b=[4,5]b=[4,5]로 바꾸고 직접 합성곱 결과와 비교해 보자.