fork download
  1. # @file convolve.002.py
  2. # @ingroup experimental
  3. # Fast convolution with packed lookup tables (GGIV).
  4. # @note Uses a lookup table and 64-bit words to calculate weights in parallel.
  5. # @date 10/05/2026
  6.  
  7. from collections import deque
  8. from itertools import islice, chain, repeat
  9. from math import floor, sumprod
  10.  
  11. def roundup(x):
  12. return floor(x + 0.5)
  13.  
  14. def partmap(func, n, iterable):
  15. iterator = iter(iterable)
  16. return chain(map(func, islice(iterator, n)), iterator)
  17.  
  18. def padtail(iterable, minsize, *, fillvalue=None):
  19. iterator = iter(iterable)
  20. iterator_with_repeat = chain(iterator, repeat(fillvalue))
  21. return chain(islice(iterator_with_repeat, minsize), iterator)
  22.  
  23. def windowed(iterable, n):
  24. iterator = iter(iterable)
  25. window = deque(islice(iterator, n-1), maxlen=n)
  26. for item in iterator:
  27. window.append(item)
  28. yield tuple(window)
  29.  
  30. def _convolve(signal, kernel):
  31. kernel = tuple(kernel)[::-1]
  32. n = len(kernel)
  33. padded_signal = chain(repeat(0, n-1), signal, repeat(0, n-1))
  34. return map(sumprod, repeat(kernel), windowed(padded_signal, n))
  35.  
  36. # Wrap _convolve to accept half-kernel and window output to input signal.
  37.  
  38. def convolve1(signal, kernel):
  39. s = tuple(signal)
  40. k = tuple(padtail(kernel, 1, fillvalue=0))
  41. n = len(k)
  42. k = k[1:][::-1]+k
  43. return islice(_convolve(k, s), n-1, n-1+len(s))
  44.  
  45. # Accepts a half-kernel whose center element is in the first position. The
  46. # remaining elements are mirrored and the kernel is padded to have a total
  47. # width of nine.
  48.  
  49. def convolve2(signal, kernel):
  50.  
  51. # Real to fixed point.
  52.  
  53. def to_i(x):
  54. return int(x * 16) # 8.4
  55.  
  56. # Pack fixed point from least to most significant position in integer.
  57.  
  58. def pack(ws):
  59. return sum(w << (i * 12) for i, w in enumerate(ws))
  60.  
  61. # Bias for kernel weights.
  62.  
  63. def getbias(w):
  64. return int(-w * (1 << 12)) if w < 0 else 0
  65.  
  66. # Halve center weight since it's counted twice.
  67.  
  68. k = partmap(lambda x: x/2, 1, kernel)
  69.  
  70. # Zero pad tail of half-kernel if needed.
  71.  
  72. k = tuple(padtail(k, 5, fillvalue=0))
  73. m = len(k)
  74.  
  75. assert m <= 5
  76. assert sum(map(abs, k)) <= 1.003, "potential overflow!"
  77.  
  78. # Zero pad signal and store original signal length.
  79.  
  80. s = tuple(chain(repeat(0, m-1), signal, repeat(0, m-1)))
  81. n = len(s) - 2*(m-1)
  82.  
  83. # Calculate bias for full-kernel (double half-kernel bias).
  84.  
  85. bias = 2*sum(map(getbias, k))
  86.  
  87. # Initialize lookup table mapping bytes to packed integers.
  88.  
  89. lut = tuple(pack(to_i(i*w) + getbias(w) for w in k) for i in range(256))
  90.  
  91. # Initialize and advance parallel sums to first writeable element.
  92.  
  93. fwd = lut[s[0]]
  94. rev = lut[s[4]]
  95.  
  96. for i in range(1, 4):
  97. fwd = ((fwd >> 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i]]
  98. rev = ((rev << 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i+4]]
  99.  
  100. # Advance and combine parallel sums for writeable elements.
  101.  
  102. for i in range(n):
  103. fwd = ((fwd >> 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i+4]]
  104. rev = ((rev << 12) & 0xFFFFFFFFFFFFFFFF) + lut[s[i+8]]
  105. val = ((fwd & 0xFFF) + ((rev >> 48) & 0xFFF) + 8 - bias) >> 4
  106. yield val
  107.  
  108. # Show.
  109.  
  110. def show(s, k):
  111. print("Signal:", s)
  112. print("Kernel:", list(map(lambda x: round(x, 6), k)))
  113. # print(list(map(lambda x: round(x, 3), convolve1(s, k))))
  114. print(list(map(roundup, convolve1(s, k))))
  115. print(list(convolve2(s, k)))
  116.  
  117. ws = []
  118. for w in (1, -3, 5, -7, 2):
  119. ws.append(w)
  120. k = [x/sum(map(abs, ws)) for x in ws]
  121. s = [1, 2, 3, 4, 5]
  122. show(s, k)
  123. s = [x * 50 for x in s]
  124. show(s, k)
Success #stdin #stdout 0.08s 14220KB
stdin
Standard input is empty
stdout
Signal: [1, 2, 3, 4, 5]
Kernel: [1.0]
[1, 2, 3, 4, 5]
[1, 2, 3, 4, 5]
Signal: [50, 100, 150, 200, 250]
Kernel: [1.0]
[50, 100, 150, 200, 250]
[50, 100, 150, 200, 250]
Signal: [1, 2, 3, 4, 5]
Kernel: [0.25, -0.75]
[-1, -2, -4, -5, -2]
[-1, -2, -4, -5, -2]
Signal: [50, 100, 150, 200, 250]
Kernel: [0.25, -0.75]
[-62, -125, -187, -250, -87]
[-62, -125, -187, -250, -87]
Signal: [1, 2, 3, 4, 5]
Kernel: [0.111111, -0.333333, 0.555556]
[1, 1, 2, -1, 1]
[1, 1, 2, -1, 1]
Signal: [50, 100, 150, 200, 250]
Kernel: [0.111111, -0.333333, 0.555556]
[56, 56, 83, -56, 44]
[56, 55, 83, -56, 44]
Signal: [1, 2, 3, 4, 5]
Kernel: [0.0625, -0.1875, 0.3125, -0.4375]
[-1, -2, 1, -1, 0]
[-1, -2, 1, -1, 0]
Signal: [50, 100, 150, 200, 250]
Kernel: [0.0625, -0.1875, 0.3125, -0.4375]
[-56, -78, 47, -53, -19]
[-56, -78, 47, -53, -19]
Signal: [1, 2, 3, 4, 5]
Kernel: [0.055556, -0.166667, 0.277778, -0.388889, 0.111111]
[0, -1, 1, -1, 0]
[0, -1, 1, -1, 0]
Signal: [50, 100, 150, 200, 250]
Kernel: [0.055556, -0.166667, 0.277778, -0.388889, 0.111111]
[-22, -69, 42, -47, -11]
[-22, -69, 42, -47, -11]