Number Theory & Algebra contents
Fast Fourier Transform and Number Theoretic Transform
Multiply polynomials and big numbers in O(n log n): the complex FFT, the exact NTT modulo 998244353, arbitrary moduli, and a big-integer shortcut that is uniquely fast in Python.
Read first: Primitive Roots, Garner's Algorithm
Multiplying two polynomials of degree by the definition takes . The Fast Fourier Transform (FFT) does it in . Since a big number is a polynomial in its base, the same technique multiplies huge integers, and many counting problems are secretly convolutions.
The idea: evaluate, multiply pointwise, interpolate
A polynomial of degree is determined by its values at distinct points. The product has degree , so evaluate and at points, multiply the values pointwise (), and interpolate back. Evaluating at arbitrary points costs , but if the points are the -th roots of unity
then divide and conquer works. The values are the discrete Fourier transform (DFT) of the coefficient vector.
Fast Fourier transform
Split into even- and odd-indexed coefficients. With even, takes only distinct values (the roots of unity of order ) and . Therefore for :
so a transform of size reduces to two of size plus work: .
import cmath
def fft(a, invert=False):
"""Recursive complex FFT; len(a) must be a power of two."""
n = len(a)
if n == 1:
return a[:]
even = fft(a[0::2], invert)
odd = fft(a[1::2], invert)
angle = 2 * cmath.pi / n * (-1 if invert else 1)
w, wn = 1, cmath.exp(1j * angle)
result = [0] * n
for k in range(n // 2):
t = w * odd[k]
result[k] = even[k] + t
result[k + n // 2] = even[k] - t
w *= wn
return result
def multiply_fft(a, b):
"""Product of two integer polynomials via floating-point FFT (rounded)."""
n = 1
while n < len(a) + len(b) - 1:
n *= 2
fa = fft([complex(x) for x in a] + [0] * (n - len(a)))
fb = fft([complex(x) for x in b] + [0] * (n - len(b)))
prod = fft([x * y for x, y in zip(fa, fb)], invert=True)
return [round(x.real / n) for x in prod[: len(a) + len(b) - 1]]
def multiply_naive(a, b):
res = [0] * (len(a) + len(b) - 1)
for i, x in enumerate(a):
for j, y in enumerate(b):
res[i + j] += x * y
return res
assert multiply_fft([1, 2, 3], [4, 5]) == [4, 13, 22, 15] # (1+2x+3x^2)(4+5x)
import random
random.seed(1)
for _ in range(100):
a = [random.randint(-50, 50) for _ in range(random.randint(1, 40))]
b = [random.randint(-50, 50) for _ in range(random.randint(1, 40))]
assert multiply_fft(a, b) == multiply_naive(a, b)The inverse transform is the same algorithm with instead of , followed by division by (the code above does the division at the end).
Number theoretic transform (NTT)
The FFT only needs a number with two properties: and for (a primitive -th root of unity). Such elements exist in modular arithmetic too. If is prime and is a primitive root, then is a primitive -th root of unity for every power of two . The transform then works over integers modulo , exactly, with no rounding.
The classic modulus: , root , allowing sizes up to .
MOD = 998244353
G = 3
def ntt(a, invert=False):
"""In-place style iterative NTT (returns a new list). len(a) must be a power of two <= 2^23."""
n = len(a)
a = a[:]
j = 0
for i in range(1, n): # bit-reversal permutation
bit = n >> 1
while j & bit:
j ^= bit
bit >>= 1
j ^= bit
if i < j:
a[i], a[j] = a[j], a[i]
length = 2
while length <= n:
w = pow(G, (MOD - 1) // length, MOD)
if invert:
w = pow(w, MOD - 2, MOD)
half = length // 2
ws = [1] * half
for k in range(1, half):
ws[k] = ws[k - 1] * w % MOD
for start in range(0, n, length):
for k in range(half):
u = a[start + k]
v = a[start + k + half] * ws[k] % MOD
a[start + k] = (u + v) % MOD
a[start + k + half] = (u - v) % MOD
length <<= 1
if invert:
n_inv = pow(n, MOD - 2, MOD)
a = [x * n_inv % MOD for x in a]
return a
def multiply_ntt(a, b):
"""Product of two polynomials modulo 998244353."""
need = len(a) + len(b) - 1
n = 1
while n < need:
n *= 2
fa = ntt(a + [0] * (n - len(a)))
fb = ntt(b + [0] * (n - len(b)))
return ntt([x * y % MOD for x, y in zip(fa, fb)], invert=True)[:need]
for _ in range(100):
a = [random.randrange(MOD) for _ in range(random.randint(1, 60))]
b = [random.randrange(MOD) for _ in range(random.randint(1, 60))]
assert multiply_ntt(a, b) == [x % MOD for x in multiply_naive(a, b)]The loop structure is the standard iterative FFT: first permute the input by reversing the bits of the indices, then combine pairs of size , , , ... in place ("butterfly" operations).
Other NTT-friendly primes with primitive root 3: and .
Multiplication modulo an arbitrary number
For a modulus like (not NTT-friendly), compute the product exactly (coefficients up to ) using three NTT primes whose product exceeds that bound, and combine the three residues with Garner's algorithm:
PRIMES = [998244353, 167772161, 469762049]
def ntt_mod(a, mod, invert=False, g=3):
n = len(a)
a = a[:]
j = 0
for i in range(1, n):
bit = n >> 1
while j & bit:
j ^= bit
bit >>= 1
j ^= bit
if i < j:
a[i], a[j] = a[j], a[i]
length = 2
while length <= n:
w = pow(g, (mod - 1) // length, mod)
if invert:
w = pow(w, mod - 2, mod)
half = length // 2
ws = [1] * half
for k in range(1, half):
ws[k] = ws[k - 1] * w % mod
for start in range(0, n, length):
for k in range(half):
u = a[start + k]
v = a[start + k + half] * ws[k] % mod
a[start + k] = (u + v) % mod
a[start + k + half] = (u - v) % mod
length <<= 1
if invert:
n_inv = pow(n, mod - 2, mod)
a = [x * n_inv % mod for x in a]
return a
def multiply_any_mod(a, b, mod):
need = len(a) + len(b) - 1
n = 1
while n < need:
n *= 2
residues = []
for p in PRIMES:
fa = ntt_mod([x % p for x in a] + [0] * (n - len(a)), p)
fb = ntt_mod([x % p for x in b] + [0] * (n - len(b)), p)
residues.append(ntt_mod([x * y % p for x, y in zip(fa, fb)], p, invert=True)[:need])
p1, p2, p3 = PRIMES
inv12 = pow(p1, -1, p2)
inv123 = pow(p1 * p2, -1, p3)
out = []
for r1, r2, r3 in zip(*residues):
x1 = r1
x2 = (r2 - x1) * inv12 % p2
x3 = ((r3 - x1 - x2 * p1) * inv123) % p3
out.append((x1 + x2 * p1 + x3 * p1 * p2) % mod) # Garner: mixed radix evaluated mod `mod`
return out
M = 10 ** 9 + 7
for _ in range(30):
a = [random.randrange(M) for _ in range(random.randint(1, 50))]
b = [random.randrange(M) for _ in range(random.randint(1, 50))]
assert multiply_any_mod(a, b, M) == [x % M for x in multiply_naive(a, b)]A Python shortcut: multiplication with big integers (Kronecker substitution)
CPython multiplies huge integers with the Karatsuba algorithm in C. So we can let Python's int do the convolution: pack the coefficients into one big integer with enough bits per coefficient to avoid overlap, multiply the integers, then unpack. If every coefficient is below and the polynomials have at most terms, each product coefficient is below , so bits per slot with suffice.
def multiply_kronecker(a, b):
"""Exact product of polynomials with non-negative integer coefficients."""
n = min(len(a), len(b))
biggest = max(max(a), max(b), 1)
bits = 2 * biggest.bit_length() + n.bit_length() + 1
nbytes = (bits + 7) // 8
A = int.from_bytes(b"".join(x.to_bytes(nbytes, "little") for x in a), "little")
B = int.from_bytes(b"".join(x.to_bytes(nbytes, "little") for x in b), "little")
C = A * B
raw = C.to_bytes(nbytes * (len(a) + len(b)), "little")
return [int.from_bytes(raw[i * nbytes : (i + 1) * nbytes], "little") for i in range(len(a) + len(b) - 1)]
for _ in range(100):
a = [random.randrange(10 ** 9) for _ in range(random.randint(1, 60))]
b = [random.randrange(10 ** 9) for _ in range(random.randint(1, 60))]
assert multiply_kronecker(a, b) == multiply_naive(a, b)We measured two random polynomials with 30-bit coefficients:
| terms | naive | pure-Python NTT (mod 998244353) | Kronecker |
|---|---|---|---|
| 1,024 | 0.21 s | 15 ms | 3 ms |
| 4,096 | (minutes) | 71 ms | 22 ms |
| 16,384 | – | 356 ms | 173 ms |
The exact integer product beats the hand-written NTT while being 6 lines long and needing no modulus at all; reduce the coefficients afterwards if needed. For negative coefficients, shift by an offset or split into positive and negative parts.
Applications
All possible sums
Given two sets of numbers, which sums are attainable? Build the indicator polynomials and ; the exponents present in the product are the attainable sums (and their coefficients count the number of ways).
def possible_sums(A, B):
pa = [0] * (max(A) + 1)
pb = [0] * (max(B) + 1)
for x in A:
pa[x] = 1
for x in B:
pb[x] = 1
prod = multiply_kronecker(pa, pb)
return {s: c for s, c in enumerate(prod) if c}
sums = possible_sums([1, 3, 4], [2, 3])
assert sums == {3: 1, 4: 1, 5: 1, 6: 2, 7: 1} # 6 = 3+3 = 4+2
assert set(sums) == {a + b for a in [1, 3, 4] for b in [2, 3]}All scalar products of cyclic shifts
To compute for every shift , reverse one array and convolve: the shifted dot products are the "correlation" of and .
def cyclic_correlation(a, b):
n = len(a)
doubled = b + b
prod = multiply_kronecker(a[::-1], doubled)
return [prod[n - 1 + s] for s in range(n)]
a, b = [1, 2, 3, 4], [5, 6, 7, 8]
expected = [sum(a[i] * b[(i + s) % 4] for i in range(4)) for s in range(4)]
assert cyclic_correlation(a, b) == expectedString matching with wildcards
Match pattern against text where the character * matches anything. Map letters to and wildcards to . Position matches exactly when
(each term is non-negative and vanishes only for a wildcard or equal letters). Expanding gives three convolutions.
def wildcard_match(text, pattern):
def code(c):
return 0 if c == "*" else ord(c) - 96
T = [code(c) for c in text]
P = [code(c) for c in pattern][::-1] # reverse: turns correlation into convolution
m = len(pattern)
t1, t2, t3 = T, [x * x for x in T], [x ** 3 for x in T]
p1, p2, p3 = P, [x * x for x in P], [x ** 3 for x in P]
c1 = multiply_kronecker(t1, p3) # T * P^3
c2 = multiply_kronecker(t2, p2) # T^2 * P^2
c3 = multiply_kronecker(t3, p1) # T^3 * P
out = []
for i in range(len(text) - m + 1):
k = i + m - 1
if c1[k] - 2 * c2[k] + c3[k] == 0:
out.append(i)
return out
assert wildcard_match("abcabcab", "ab*") == [0, 3]
assert wildcard_match("a*c", "abc") == [0]
def wildcard_naive(text, pattern):
return [i for i in range(len(text) - len(pattern) + 1)
if all(p == "*" or t == "*" or p == t for p, t in zip(pattern, text[i:i + len(pattern)]))]
for _ in range(200):
text = "".join(random.choice("ab*") for _ in range(random.randint(1, 20)))
pat = "".join(random.choice("ab*") for _ in range(random.randint(1, 5)))
assert wildcard_match(text, pat) == wildcard_naive(text, pat)Choosing a method
| Task | Best tool in Python |
|---|---|
| product of integer polynomials, exact | Kronecker substitution (big-int *) |
| product modulo 998244353 | NTT, or Kronecker followed by % MOD |
| product modulo other | Kronecker followed by % m |
| multiplication of huge integers | just use * (Karatsuba built in) |
| you need the transform itself (for pointwise operations, polynomial inverse, ...) | NTT |
Practice problems
- SPOJ - POLYMUL
- SPOJ - MAXMATCH
- SPOJ - ADAMATCH
- Codeforces - Yet Another String Matching Problem
- Codeforces - Lightsabers (hard)
- Codeforces - Running Competition
- Kattis - A+B Problem
- Kattis - K-Inversions
- Codeforces - Dasha and cyclic table
- CodeChef - Expected Number of Customers
- CodeChef - Power Sum
- Codeforces - Centroid Probabilities