1"""A codec that beats LPC+ANS on the quantized filtered-Gaussian process.
3The insight: with prediction-error std ~0.28 quantization steps, the
4conditional distribution of z_t is a Gaussian bump over 1-3 integer bins whose
5placement depends on the *fractional part* of the real-valued prediction.
6Integer-residual LPC + a memoryless entropy coder pools all fractional parts
7into one histogram and pays the mixture entropy (~1.10 bits). Coding z_t
8against a discretized Gaussian centred at the real-valued prediction pays the
9conditional entropy (~0.72 bits) instead.
11Codec = order-p linear predictor (coefficients fitted on the block, sent as
12float32 in the header) + adaptive binary arithmetic coding of z_t under
13N(mu_t, s^2) discretized to integer bins, where mu_t is the float prediction
14from already-decoded samples and s is a single fitted scale sent in the
15header. Encoder and decoder compute mu_t with the identical np.dot on
16identical float64 data, so the frequency tables agree bit for bit.
18Everything the decoder needs is charged: order, coefficients, s, marginal
19std for the warm-up samples, and the sample count.
20"""
21import math
22import numpy as np
23from scipy.linalg import solve_toeplitz
24from scipy.signal import lfilter
25from scipy.special import ndtr
27SIGMA = 5.0
28RATE = 30000.0
29LOW, HIGH, TAPS = 300.0, 2000.0, 31
31TOTAL_BITS = 16
32TOTAL = 1 << TOTAL_BITS
33KWIN = 32 # symbols are residuals in [-KWIN, KWIN-1] around round(mu)
34SQRT2 = math.sqrt(2.0)
37# ---------------------------------------------------------------- the process
38def windowed_sinc_lowpass(fc, taps):
39 n = taps | 1
40 mid = (n - 1) / 2
41 i = np.arange(n)
42 t = i - mid
43 sinc = np.where(t == 0, 2 * fc,
44 np.sin(2 * np.pi * fc * t) / (np.pi * np.where(t == 0, 1, t)))
45 w = 0.54 - 0.46 * np.cos(2 * np.pi * i / (n - 1))
46 h = sinc * w
47 return h / h.sum()
50def make_signal(n, seed=1):
51 h = windowed_sinc_lowpass(HIGH / RATE, TAPS) - windowed_sinc_lowpass(LOW / RATE, TAPS)
52 rng = np.random.default_rng(seed)
53 x = SIGMA * rng.standard_normal(n + len(h) - 1)
54 y = np.convolve(x, h, mode='valid')
55 return np.floor(y + 0.5).astype(np.int64)
58# ------------------------------------------------- arithmetic coder (WNC-style)
59class BitWriter:
60 def __init__(self):
61 self.bytes = bytearray()
62 self.acc = 0
63 self.nbits = 0
65 def write(self, bit):
66 self.acc = (self.acc << 1) | bit
67 self.nbits += 1
68 if self.nbits == 8:
69 self.bytes.append(self.acc)
70 self.acc = 0
71 self.nbits = 0
73 def flush(self):
74 while self.nbits:
75 self.write(0)
78class BitReader:
79 def __init__(self, data):
80 self.data = data
81 self.pos = 0
83 def read(self):
84 byte = self.data[self.pos >> 3] if (self.pos >> 3) < len(self.data) else 0
85 bit = (byte >> (7 - (self.pos & 7))) & 1
86 self.pos += 1
87 return bit
90class ArithEncoder:
91 FULL = (1 << 32) - 1
92 HALF = 1 << 31
93 QUARTER = 1 << 30
95 def __init__(self):
96 self.low = 0
97 self.high = self.FULL
98 self.pending = 0
99 self.out = BitWriter()
101 def _emit(self, bit):
102 self.out.write(bit)
103 while self.pending:
104 self.out.write(1 - bit)
105 self.pending -= 1
107 def encode(self, cum, freq, tot):
108 span = self.high - self.low + 1
109 self.high = self.low + span * (cum + freq) // tot - 1
110 self.low = self.low + span * cum // tot
111 while True:
112 if self.high < self.HALF:
113 self._emit(0)
114 elif self.low >= self.HALF:
115 self._emit(1)
116 self.low -= self.HALF
117 self.high -= self.HALF
118 elif self.low >= self.QUARTER and self.high < self.HALF + self.QUARTER:
119 self.pending += 1
120 self.low -= self.QUARTER
121 self.high -= self.QUARTER
122 else:
123 break
124 self.low <<= 1
125 self.high = (self.high << 1) | 1
127 def finish(self):
128 self.pending += 1
129 self._emit(0 if self.low < self.QUARTER else 1)
130 self.out.flush()
131 return bytes(self.out.bytes)
134class ArithDecoder:
135 FULL = (1 << 32) - 1
136 HALF = 1 << 31
137 QUARTER = 1 << 30
139 def __init__(self, data):
140 self.low = 0
141 self.high = self.FULL
142 self.inp = BitReader(data)
143 self.value = 0
144 for _ in range(32):
145 self.value = (self.value << 1) | self.inp.read()
147 def target(self, tot):
148 span = self.high - self.low + 1
149 return ((self.value - self.low + 1) * tot - 1) // span
151 def consume(self, cum, freq, tot):
152 span = self.high - self.low + 1
153 self.high = self.low + span * (cum + freq) // tot - 1
154 self.low = self.low + span * cum // tot
155 while True:
156 if self.high < self.HALF:
157 pass
158 elif self.low >= self.HALF:
159 self.low -= self.HALF
160 self.high -= self.HALF
161 self.value -= self.HALF
162 elif self.low >= self.QUARTER and self.high < self.HALF + self.QUARTER:
163 self.low -= self.QUARTER
164 self.high -= self.QUARTER
165 self.value -= self.QUARTER
166 else:
167 break
168 self.low <<= 1
169 self.high = (self.high << 1) | 1
170 self.value = (self.value << 1) | self.inp.read()
173# --------------------------------------------- discretized-Gaussian frequencies
174def freq_table(d, s):
175 """Integer frequencies (sum TOTAL) for residual symbols -KWIN..KWIN-1:
176 bin k has probability P(round(N(d, s^2)) = k), tails folded into the edge
177 bins, every bin floored at 1. Deterministic in (d, s)."""
178 inv = 1.0 / (s * SQRT2)
179 m = min(KWIN - 1, int(8.0 * s) + 2)
180 freqs = [1] * (2 * KWIN)
181 spare = TOTAL - 2 * KWIN
182 lo_cdf = 0.0
183 probs = []
184 for k in range(-m, m + 1):
185 hi_cdf = 1.0 if k == m else 0.5 * math.erfc(-(k + 0.5 - d) * inv)
186 probs.append(hi_cdf - lo_cdf)
187 lo_cdf = hi_cdf
188 scaled = [int(p * spare) for p in probs]
189 deficit = spare - sum(scaled)
190 scaled[max(range(len(scaled)), key=scaled.__getitem__)] += deficit
191 for i, k in enumerate(range(-m, m + 1)):
192 freqs[k + KWIN] += scaled[i]
193 return freqs
196def cumulative(freqs):
197 cum = [0] * (len(freqs) + 1)
198 for i, f in enumerate(freqs):
199 cum[i + 1] = cum[i] + f
200 return cum
203# ------------------------------------------------------------------- the codec
204def fit_model(z, order):
205 zf = z.astype(np.float64)
206 r = np.array([zf @ zf if lag == 0 else zf[lag:] @ zf[:-lag]
207 for lag in range(order + 1)]) / len(zf)
208 a = solve_toeplitz(r[:order], r[1:order + 1]).astype(np.float32)
209 # residual scale: fit s by minimizing the ideal code length on the block
210 pred = lfilter(np.concatenate(([0.0], a.astype(np.float64))), [1.0], zf)
211 e = zf[order:] - pred[order:]
212 s0 = math.sqrt(max(e.var() - 1.0 / 12.0, 1e-6))
213 best = (None, np.inf)
214 zt, mu = zf[order:], pred[order:]
215 for s in s0 * np.linspace(0.85, 1.15, 13):
216 p = ndtr((zt + 0.5 - mu) / s) - ndtr((zt - 0.5 - mu) / s)
217 bits = float(-np.log2(np.maximum(p, 1e-12)).mean())
218 if bits < best[1]:
219 best = (s, bits)
220 s = np.float32(best[0])
221 std_z = np.float32(max(zf.std(), 1e-3))
222 return a, s, std_z, best[1]
225HEADER_BYTES = 4 + 2 + 4 + 4 # n, order, s, std_z (+ coefficients, counted below)
228def encode(z, order):
229 a, s, std_z, ideal_bits = fit_model(z, order)
230 arev = a[::-1].astype(np.float64)
231 zf = z.astype(np.float64)
232 enc = ArithEncoder()
234 warm = freq_table(0.0, float(std_z))
235 warm_cum = cumulative(warm)
236 s_f = float(s)
237 for t in range(len(z)):
238 if t < order:
239 table, cum = warm, warm_cum
240 c = 0
241 else:
242 mu = float(np.dot(arev, zf[t - order:t]))
243 c = math.floor(mu + 0.5)
244 table = freq_table(mu - c, s_f)
245 cum = cumulative(table)
246 sym = int(z[t]) - c + KWIN
247 if not 0 <= sym < 2 * KWIN:
248 raise ValueError(f'residual out of range at {t}') # production: escape code
249 enc.encode(cum[sym], table[sym], TOTAL)
250 payload = enc.finish()
251 header = HEADER_BYTES + 4 * order
252 return payload, header, a, s, std_z, ideal_bits
255def decode(payload, n, order, a, s, std_z):
256 arev = a[::-1].astype(np.float64)
257 zf = np.zeros(n, dtype=np.float64)
258 z = np.zeros(n, dtype=np.int64)
259 dec = ArithDecoder(payload)
260 warm = freq_table(0.0, float(std_z))
261 warm_cum = cumulative(warm)
262 s_f = float(s)
263 for t in range(n):
264 if t < order:
265 table, cum = warm, warm_cum
266 c = 0
267 else:
268 mu = float(np.dot(arev, zf[t - order:t]))
269 c = math.floor(mu + 0.5)
270 table = freq_table(mu - c, s_f)
271 cum = cumulative(table)
272 tgt = dec.target(TOTAL)
273 # find symbol: cum[sym] <= tgt < cum[sym+1]
274 lo, hi = 0, 2 * KWIN
275 while hi - lo > 1:
276 mid = (lo + hi) // 2
277 if cum[mid] <= tgt:
278 lo = mid
279 else:
280 hi = mid
281 sym = lo
282 dec.consume(cum[sym], table[sym], TOTAL)
283 z[t] = sym - KWIN + c
284 zf[t] = float(z[t])
285 return z
288def main():
289 import time
290 n = 1 << 20
291 order = 64
292 z = make_signal(n)
294 t0 = time.time()
295 payload, header, a, s, std_z, ideal_bits = encode(z, order)
296 t1 = time.time()
297 total_bytes = len(payload) + header
298 bps = 8.0 * total_bytes / n
299 print(f'encoded {n} samples in {t1 - t0:.1f}s')
300 print(f' ideal model cross-entropy (block fit): {ideal_bits:.4f} bits/sample')
301 print(f' payload {len(payload)} B + header {header} B = {total_bytes} B')
302 print(f' -> {bps:.4f} bits/sample ratio vs int16: {16 / bps:.2f}x')
304 t0 = time.time()
305 zdec = decode(payload, n, order, a, s, std_z)
306 t1 = time.time()
307 ok = bool(np.array_equal(z, zdec))
308 print(f'decoded in {t1 - t0:.1f}s round-trip exact: {ok}')
309 if not ok:
310 bad = np.nonzero(z != zdec)[0][:5]
311 print(f' first mismatches at {bad}')
314if __name__ == '__main__':
315 main()