import sage.all as sage
from sage.all import *
import math
import numpy as np
from typing import List, Tuple, Optional
class FFTNTTMultiplication:
def __init__(self, ntt_modulus: Optional[int] = None):
self.ntt_modulus = ntt_modulus
if ntt_modulus:
self.ntt_params = self._find_ntt_parameters(ntt_modulus)
def _find_ntt_parameters(self, modulus: int) -> dict:
for g in range(2, modulus):
if pow(g, modulus-1, modulus) == 1:
is_primitive = True
for i in range(2, modulus-1):
if pow(g, i, modulus) == 1:
is_primitive = False
break
if is_primitive:
primitive_root = g
break
else:
raise ValueError(f"No primitive root found for modulus {modulus}")
max_length = 1
while (modulus - 1) % (2 * max_length) == 0:
max_length *= 2
return {
'modulus': modulus,
'primitive_root': primitive_root,
'max_length': max_length
}
def cooley_tukey_fft(self, coeffs: List[complex], inverse: bool = False) -> List[complex]:
n = len(coeffs)
if n <= 1:
return coeffs
if n & (n - 1) != 0:
raise ValueError("FFT length must be power of 2")
even = self.cooley_tukey_fft(coeffs[::2], inverse)
odd = self.cooley_tukey_fft(coeffs[1::2], inverse)
result = [0] * n
angle = (-2 if inverse else 2) * pi * I / n
for k in range(n//2):
t = exp(angle * k) * odd[k]
result[k] = even[k] + t
result[k + n//2] = even[k] - t
return result
def number_theoretic_transform(self, coeffs: List[int], inverse: bool = False) -> List[int]:
if not self.ntt_modulus:
raise ValueError("NTT modulus not set")
n = len(coeffs)
if n <= 1:
return coeffs
if n & (n - 1) != 0:
raise ValueError("NTT length must be power of 2")
params = self.ntt_params
modulus = params['modulus']
primitive_root = params['primitive_root']
if inverse:
root = pow(primitive_root, (modulus - 1) // n, modulus)
root = pow(root, modulus - 2, modulus) else:
root = pow(primitive_root, (modulus - 1) // n, modulus)
def ntt_recursive(coeffs, root):
n = len(coeffs)
if n == 1:
return coeffs
even = ntt_recursive(coeffs[::2], pow(root, 2, modulus))
odd = ntt_recursive(coeffs[1::2], pow(root, 2, modulus))
result = [0] * n
current_root = 1
for k in range(n//2):
t = (current_root * odd[k]) % modulus
result[k] = (even[k] + t) % modulus
result[k + n//2] = (even[k] - t) % modulus
current_root = (current_root * root) % modulus
return result
result = ntt_recursive(coeffs, root)
if inverse:
n_inv = pow(n, modulus - 2, modulus) result = [(x * n_inv) % modulus for x in result]
return result
def polynomial_multiply_fft(self, a: List[complex], b: List[complex]) -> List[complex]:
n = 1
while n < len(a) + len(b) - 1:
n *= 2
a_padded = a + [0] * (n - len(a))
b_padded = b + [0] * (n - len(b))
a_fft = self.cooley_tukey_fft(a_padded)
b_fft = self.cooley_tukey_fft(b_padded)
c_fft = [x * y for x, y in zip(a_fft, b_fft)]
c = self.cooley_tukey_fft(c_fft, inverse=True)
result = [abs(x)/n for x in c]
return [round(x) for x in result]
def polynomial_multiply_ntt(self, a: List[int], b: List[int]) -> List[int]:
if not self.ntt_modulus:
raise ValueError("NTT modulus not set")
n = 1
while n < len(a) + len(b) - 1:
n *= 2
a_padded = a + [0] * (n - len(a))
b_padded = b + [0] * (n - len(b))
a_ntt = self.number_theoretic_transform(a_padded)
b_ntt = self.number_theoretic_transform(b_padded)
modulus = self.ntt_modulus
c_ntt = [(x * y) % modulus for x, y in zip(a_ntt, b_ntt)]
c = self.number_theoretic_transform(c_ntt, inverse=True)
return c
def big_integer_multiply_fft(self, a: int, b: int) -> int:
base = 10**9
def to_digits(n):
if n == 0:
return [0]
digits = []
while n > 0:
digits.append(n % base)
n //= base
return digits
a_digits = to_digits(a)
b_digits = to_digits(b)
product_coeffs = self.polynomial_multiply_fft(a_digits, b_digits)
result = 0
for i, coeff in enumerate(product_coeffs):
result += int(coeff) * (base ** i)
return result
def big_integer_multiply_ntt(self, a: int, b: int) -> int:
if not self.ntt_modulus:
raise ValueError("NTT modulus not set")
modulus = self.ntt_modulus
base = modulus // 10
def to_digits(n):
if n == 0:
return [0]
digits = []
while n > 0:
digits.append(n % base)
n //= base
return digits
a_digits = to_digits(a)
b_digits = to_digits(b)
product_coeffs = self.polynomial_multiply_ntt(a_digits, b_digits)
result = 0
for i, coeff in enumerate(product_coeffs):
result += coeff * (base ** i)
return result
def multiply(self, a: int, b: int, method: str = 'ntt') -> int:
if method == 'fft':
return self.big_integer_multiply_fft(a, b)
elif method == 'ntt':
return self.big_integer_multiply_ntt(a, b)
else:
raise ValueError(f"Unknown method: {method}")
def test_correctness(self):
print("Testing FFT/NTT multiplication correctness...")
test_cases = [
(0, 0),
(1, 1),
(123, 456),
(1234, 5678),
(10**10, 10**10),
(2**100, 2**100),
]
print("\nTesting FFT:")
for a, b in test_cases:
try:
expected = a * b
result = self.big_integer_multiply_fft(a, b)
correct = (result == expected)
print(f" {a} * {b}: {'✓' if correct else '✗'}")
if not correct:
print(f" Expected: {expected}")
print(f" Got: {result}")
except Exception as e:
print(f" {a} * {b}: Error - {e}")
print("\nTesting NTT:")
ntt_moduli = [13, 17, 97, 193, 257, 7681]
for modulus in ntt_moduli:
try:
self.ntt_modulus = modulus
self.ntt_params = self._find_ntt_parameters(modulus)
print(f"\n Modulus {modulus}:")
for a, b in test_cases[:3]: expected = a * b
result = self.big_integer_multiply_ntt(a, b)
correct = (result == expected)
print(f" {a} * {b}: {'✓' if correct else '✗'}")
except Exception as e:
print(f" Modulus {modulus}: Error - {e}")
def benchmark(self, sizes=None):
if sizes is None:
sizes = [10, 100, 1000, 10000]
print("\nBenchmarking FFT/NTT multiplication...")
print("Digits | FFT Time (s) | NTT Time (s) | Standard (s)")
print("-" * 50)
import time
for digits in sizes:
a = 10**digits - 1 b = 10**(digits//2) - 1
try:
start = time.time()
fft_result = self.big_integer_multiply_fft(a, b)
fft_time = time.time() - start
except Exception as e:
fft_time = float('inf')
print(f"FFT error: {e}")
try:
self.ntt_modulus = 998244353 self.ntt_params = self._find_ntt_parameters(self.ntt_modulus)
start = time.time()
ntt_result = self.big_integer_multiply_ntt(a, b)
ntt_time = time.time() - start
except Exception as e:
ntt_time = float('inf')
print(f"NTT error: {e}")
try:
start = time.time()
std_result = a * b
std_time = time.time() - start
except Exception as e:
std_time = float('inf')
if digits <= 100:
expected = a * b
if 'fft_result' in locals() and fft_result != expected:
print(f" FFT gave wrong result for {digits} digits!")
if 'ntt_result' in locals() and ntt_result != expected:
print(f" NTT gave wrong result for {digits} digits!")
print(">6"
">12.3f"
">12.3f"
">12.3f")
if __name__ == "__main__":
ntt_modulus = 998244353 mult = FFTNTTMultiplication(ntt_modulus)
mult.test_correctness()
mult.benchmark([10, 50, 100])