Chapter 14

합성곱

같은 차수의 계수 모으기
핵심 질문

같은 차수로 모이는 계수들은 어떻게 계산할까?

합성곱

핵심 질문

다항식 곱셈에서 같은 차수로 모이는 계수들은 어떻게 계산할까?

교과서 설명

이 자료에서는 합성곱을 먼저 다항식 계수 조합으로 본다. 두 계수 배열 , 를 곱해서 새 배열 를 만들 때, 는 차수가 가 되는 모든 곱을 더한 값이다.

같은 식을 로 바꾸어 쓰면 다음과 같다.

신호 처리에서는 이 계산을 “한 배열을 뒤집고 밀어 가며 겹침을 더한다”고 설명하기도 한다. 하지만 FFT와 다항식 곱셈을 연결할 때는 “인 계수 조합을 모은다”는 관점이 더 직접적이다.

길이가 인 배열과 길이가 인 배열을 선형 합성곱하면 결과 길이는 이다. 없는 칸은 이라고 생각한다.

결과 길이가 인 이유는 가능한 차수 합의 범위 때문이다. 가장 작은 합은 이고, 가장 큰 합은

이다. 따라서 결과 인덱스는

이고, 칸의 개수는 개다.

합성곱 공식의 유도

두 다항식을

라고 쓰자. 곱하면

이고, 분배법칙으로 모든 항을 곱하면

가 된다. 이제 결과 다항식을

라고 쓰면, 의 계수 가 되는 모든 항에서 나온다. 따라서

가 된다. 이것이 합성곱 공식이다.

직접 계산하면 모든 쌍 를 확인해야 하므로 대략 번의 곱셈이 필요하다. 두 배열 길이가 모두 정도라면 계산이 된다. FFT를 쓰는 이유는 이 병목을 줄이기 위해서다.

또 하나 구분해야 할 말이 있다. 지금 정의한 것은 선형 합성곱이다. DFT 안에서 길이를 고정한 채 곱하면 결과가 끝에서 앞으로 말려 들어가는 순환 합성곱이 된다. 그래서 FFT로 선형 합성곱을 하려면 뒤에서 zero padding을 반드시 다룬다.

직관 비유

두 주사위를 던져 합이 가 되는 경우를 모은다고 생각하자. 합이 인 경우는 처럼 여러 조합이 있다. 합성곱도 차수의 합이 가 되는 곱들을 모두 모아 더한다.

예제

가운데 계수 에서 나온다.

배열로 쓰면 같은 계산이다.

각 칸은 다음처럼 채워진다.

차수 조합표로 쓰면 다음과 같다.

이 표에서 는 결과 배열의 번호이자 다항식의 차수다. 같은 앞에 모이는 곱을 모두 더하면 가 된다.

손풀이 체크

  1. 의 가운데 값은?

    답 보기

    23+11=72\cdot3+1\cdot1=7

  2. 길이 배열과 길이 배열의 선형 합성곱 결과 길이는?

    답 보기

    3+41=63+4-1=6

  3. 의 계수 배열은?

    답 보기

    [1,2,2,1][1,2,2,1]

다음으로 이어지는 생각

그냥 계산하면 느리지만, DFT와 FFT를 이용하면 빠르게 만들 수 있다. 다음 장에서는 “느리다”와 “빠르다”를 계산량으로 비교한다.

이번 장에서 기억할 3문장

  1. 선형 합성곱은 i+j=ki+j=k가 되는 모든 계수 곱을 c[k]c[k]에 모으는 계산이다.
  2. 길이 mmnn의 선형 합성곱 결과 길이는 m+n1m+n-1이다.
  3. 직접 계산은 많은 쌍을 확인해야 하므로 큰 입력에서는 FFT가 필요해진다.

C++ Practice

C++로 확인하기

다항식 곱셈 관점으로 합성곱 c[k]=i+j=ka[i]b[j]c[k]=\sum_{i+j=k}a[i]b[j] 계산하기

#include <iostream>
#include <vector>

using namespace std;

vector<int> convolution(const vector<int>& a, const vector<int>& b) {
    vector<int> c(a.size() + b.size() - 1, 0);
    for (int i = 0; i < static_cast<int>(a.size()); ++i) {
        for (int j = 0; j < static_cast<int>(b.size()); ++j) {
            c[i + j] += a[i] * b[j];
        }
    }
    return c;
}

int main() {
    vector<int> a = {1, 2};
    vector<int> b = {3, 4};
    auto c = convolution(a, b);

    for (int value : c) {
        cout << value << ' ';
    }
    cout << "\n";
}

연습: aa, bb를 더 긴 배열로 바꾸고 결과 길이가 m+n1m+n-1인지 확인해 보자.