# @file convolve.002.py
# @ingroup experimental
# Fast convolution with packed lookup tables (GGIV).
# @note Uses a lookup table and 64-bit words to calculate weights in parallel.
# @date 10/05/2026

from collections import deque
from itertools import islice, chain, repeat
from math import floor, sumprod

def roundup(x):
    return floor(x + 0.5)

def partmap(func, n, iterable):
    iterator = iter(iterable)
    return chain(map(func, islice(iterator, n)), iterator)

def padtail(iterable, minsize, *, fillvalue=None):
    iterator = iter(iterable)
    iterator_with_repeat = chain(iterator, repeat(fillvalue))
    return chain(islice(iterator_with_repeat, minsize), iterator)

def windowed(iterable, n):
    iterator = iter(iterable)
    window = deque(islice(iterator, n-1), maxlen=n)
    for item in iterator:
        window.append(item)
        yield tuple(window)

def _convolve(signal, kernel):
    kernel = tuple(kernel)[::-1]
    n = len(kernel)
    padded_signal = chain(repeat(0, n-1), signal, repeat(0, n-1))
    return map(sumprod, repeat(kernel), windowed(padded_signal, n))

# Wrap _convolve to accept half-kernel and window output to input signal.

def convolve1(signal, kernel):
    s = tuple(signal)
    k = tuple(padtail(kernel, 1, fillvalue=0))
    n = len(k)
    k = k[1:][::-1]+k
    return islice(_convolve(k, s), n-1, n-1+len(s))

# Accepts a half-kernel whose center element is in the first position. The
# remaining elements are mirrored and the kernel is padded to have a total
# width of nine.

def convolve2(signal, kernel):

    # Real to fixed point.

    def to_i(x):
        return int(x * 16) # 8.4

    # Pack fixed point from least to most significant position in integer.

    def pack(ws):
        return sum(w << (i * 12) for i, w in enumerate(ws))

    # Bias for kernel weights.

    def getbias(w):
        return int(-w * (1 << 12)) if w < 0 else 0

    # Halve center weight since it's counted twice.

    k = partmap(lambda x: x/2, 1, kernel)

    # Zero pad tail of half-kernel if needed.

    k = tuple(padtail(k, 5, fillvalue=0))
    m = len(k)

    assert m <= 5
    assert sum(map(abs, k)) <= 1.003, "potential overflow!"

    # Zero pad signal and store original signal length.

    s = tuple(chain(repeat(0, m-1), signal, repeat(0, m-1)))
    n = len(s) - 2*(m-1)

    # Calculate bias for full-kernel (double half-kernel bias).

    bias = 2*sum(map(getbias, k))

    # Initialize lookup table mapping bytes to packed integers.

    lut = tuple(pack(to_i(i*w) + getbias(w) for w in k) for i in range(256))

    # Initialize and advance parallel sums to first writeable element.

    fwd = lut[s[0]]
    rev = lut[s[4]]

    for i in range(1, 4):
        fwd = ((fwd >> 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i]]
        rev = ((rev << 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i+4]]

    # Advance and combine parallel sums for writeable elements.

    for i in range(n):
        fwd = ((fwd >> 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i+4]]
        rev = ((rev << 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i+8]]
        val = ((fwd & 0xFFF) + ((rev >> 48) & 0xFFF) + 8 - bias) >> 4
        yield val

# Show.

def show(s, k):
    print("Signal:", s)
    print("Kernel:", list(map(lambda x: round(x, 6), k)))
    # print(list(map(lambda x: round(x, 3), convolve1(s, k))))
    print(list(map(roundup, convolve1(s, k))))
    print(list(convolve2(s, k)))

ws = []
for w in (1, -3, 5, -7, 2):
    ws.append(w)
    k = [x/sum(map(abs, ws)) for x in ws]
    s = [1, 2, 3, 4, 5]
    show(s, k)
    s = [x * 50 for x in s]
    show(s, k)