polypress 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
polypress/codec.py ADDED
@@ -0,0 +1,343 @@
1
+ """Lossless table codec built on low-order polynomial extrapolation.
2
+
3
+ The salvaged form of the idea. Per column:
4
+
5
+ 1. Read the printed cells as fixed-point integers (scientific tables are
6
+ printed to a fixed number of decimals, so this is exact).
7
+ 2. Predict each value by extrapolating a degree-k polynomial fitted to the
8
+ k+1 values before it. On the integer grid that prediction is just a
9
+ finite-difference formula, so the residual is the (k+1)-th finite
10
+ difference -- the discrete analogue of a Taylor remainder.
11
+ 3. Entropy-code the residuals with Rice coding.
12
+
13
+ Nothing is thrown away: the residual stream restores the column exactly. The
14
+ saving comes from the residuals being small, which happens exactly when the
15
+ column is locally well-approximated by a low-degree polynomial.
16
+
17
+ Round-trips to the original CSV text byte-for-byte.
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ import re
23
+ from typing import List, Optional, Sequence, Tuple
24
+
25
+ BLOCK = 64 # residuals per Rice-parameter block
26
+ MAX_ORDER = 4 # highest polynomial predictor order
27
+ RICE_ESCAPE = 48 # unary quotient cap before falling back to raw varint
28
+ K_BITS = 6 # width of the per-block Rice-parameter field
29
+ MAX_K = (1 << K_BITS) - 1
30
+
31
+ # pred[i] = sum(c[j] * x[i-1-j]); residual = x[i] - pred[i].
32
+ # These are the degree-k polynomial extrapolators on a unit grid.
33
+ FIXED_PREDICTORS = {
34
+ 0: [],
35
+ 1: [1],
36
+ 2: [2, -1],
37
+ 3: [3, -3, 1],
38
+ 4: [4, -6, 4, -1],
39
+ }
40
+
41
+ _NUMERIC = re.compile(r"^-?\d+(\.\d+)?$")
42
+
43
+
44
+ class BitWriter:
45
+ def __init__(self) -> None:
46
+ self.buf = bytearray()
47
+ self._cur = 0
48
+ self._nbits = 0
49
+
50
+ def bit(self, b: int) -> None:
51
+ self._cur = (self._cur << 1) | (b & 1)
52
+ self._nbits += 1
53
+ if self._nbits == 8:
54
+ self.buf.append(self._cur)
55
+ self._cur = 0
56
+ self._nbits = 0
57
+
58
+ def bits(self, value: int, width: int) -> None:
59
+ for shift in range(width - 1, -1, -1):
60
+ self.bit((value >> shift) & 1)
61
+
62
+ def unary(self, q: int) -> None:
63
+ for _ in range(q):
64
+ self.bit(1)
65
+ self.bit(0)
66
+
67
+ def uvarint(self, v: int) -> None:
68
+ assert v >= 0
69
+ while True:
70
+ chunk = v & 0x7F
71
+ v >>= 7
72
+ self.bit(1 if v else 0)
73
+ self.bits(chunk, 7)
74
+ if not v:
75
+ return
76
+
77
+ def svarint(self, v: int) -> None:
78
+ self.uvarint(zigzag(v))
79
+
80
+ def rice(self, u: int, k: int) -> None:
81
+ q = u >> k
82
+ if q >= RICE_ESCAPE:
83
+ self.unary(RICE_ESCAPE)
84
+ self.uvarint(u)
85
+ else:
86
+ self.unary(q)
87
+ if k:
88
+ self.bits(u & ((1 << k) - 1), k)
89
+
90
+ def bytes_out(self) -> bytes:
91
+ if self._nbits:
92
+ self.buf.append(self._cur << (8 - self._nbits))
93
+ self._cur = 0
94
+ self._nbits = 0
95
+ return bytes(self.buf)
96
+
97
+
98
+ class BitReader:
99
+ def __init__(self, data: bytes) -> None:
100
+ self.data = data
101
+ self.pos = 0
102
+
103
+ def bit(self) -> int:
104
+ byte = self.data[self.pos >> 3]
105
+ b = (byte >> (7 - (self.pos & 7))) & 1
106
+ self.pos += 1
107
+ return b
108
+
109
+ def bits(self, width: int) -> int:
110
+ v = 0
111
+ for _ in range(width):
112
+ v = (v << 1) | self.bit()
113
+ return v
114
+
115
+ def unary(self) -> int:
116
+ q = 0
117
+ while self.bit():
118
+ q += 1
119
+ return q
120
+
121
+ def uvarint(self) -> int:
122
+ v = 0
123
+ shift = 0
124
+ while True:
125
+ more = self.bit()
126
+ v |= self.bits(7) << shift
127
+ shift += 7
128
+ if not more:
129
+ return v
130
+
131
+ def svarint(self) -> int:
132
+ return unzigzag(self.uvarint())
133
+
134
+ def rice(self, k: int) -> int:
135
+ q = self.unary()
136
+ if q >= RICE_ESCAPE:
137
+ return self.uvarint()
138
+ return (q << k) | (self.bits(k) if k else 0)
139
+
140
+
141
+ def zigzag(v: int) -> int:
142
+ return 2 * v if v >= 0 else -2 * v - 1
143
+
144
+
145
+ def unzigzag(u: int) -> int:
146
+ return u >> 1 if u % 2 == 0 else -((u + 1) >> 1)
147
+
148
+
149
+ # ---------------------------------------------------------------- fixed point
150
+
151
+ def cell_to_int(cell: str, decimals: int) -> int:
152
+ neg = cell.startswith("-")
153
+ if neg:
154
+ cell = cell[1:]
155
+ if "." in cell:
156
+ whole, frac = cell.split(".")
157
+ else:
158
+ whole, frac = cell, ""
159
+ frac = (frac + "0" * decimals)[:decimals]
160
+ v = int(whole + frac)
161
+ return -v if neg else v
162
+
163
+
164
+ def int_to_cell(value: int, decimals: int) -> str:
165
+ neg = value < 0
166
+ digits = str(abs(value)).rjust(decimals + 1, "0")
167
+ if decimals:
168
+ digits = digits[:-decimals] + "." + digits[-decimals:]
169
+ return ("-" if neg else "") + digits
170
+
171
+
172
+ def as_numeric_column(cells: Sequence[str]) -> Optional[Tuple[List[int], int]]:
173
+ """Return (integers, decimals) if the column can be stored as fixed point
174
+ and reproduced exactly, else None."""
175
+ decimals = 0
176
+ for cell in cells:
177
+ if not _NUMERIC.match(cell):
178
+ return None
179
+ if "." in cell:
180
+ decimals = max(decimals, len(cell.split(".")[1]))
181
+ if decimals > 63:
182
+ return None
183
+ ints = [cell_to_int(c, decimals) for c in cells]
184
+ for value, original in zip(ints, cells):
185
+ if int_to_cell(value, decimals) != original:
186
+ return None # leading zeros, "-0.0", ragged decimals, ...
187
+ return ints, decimals
188
+
189
+
190
+ # ------------------------------------------------------------- prediction
191
+
192
+ def residuals(ints: Sequence[int], order: int) -> List[int]:
193
+ coeffs = FIXED_PREDICTORS[order]
194
+ out = []
195
+ for i in range(order, len(ints)):
196
+ pred = 0
197
+ for j, c in enumerate(coeffs):
198
+ pred += c * ints[i - 1 - j]
199
+ out.append(ints[i] - pred)
200
+ return out
201
+
202
+
203
+ def restore(warmup: Sequence[int], res: Sequence[int], order: int, n: int) -> List[int]:
204
+ coeffs = FIXED_PREDICTORS[order]
205
+ ints = list(warmup)
206
+ for i in range(order, n):
207
+ pred = 0
208
+ for j, c in enumerate(coeffs):
209
+ pred += c * ints[i - 1 - j]
210
+ ints.append(res[i - order] + pred)
211
+ return ints
212
+
213
+
214
+ def _rice_cost(us: Sequence[int], k: int) -> int:
215
+ total = 0
216
+ for u in us:
217
+ q = u >> k
218
+ if q >= RICE_ESCAPE:
219
+ nbytes = max(1, (u.bit_length() + 6) // 7)
220
+ total += RICE_ESCAPE + 1 + 8 * nbytes
221
+ else:
222
+ total += q + 1 + k
223
+ return total
224
+
225
+
226
+ def _best_k(us: Sequence[int]) -> Tuple[int, int]:
227
+ best = (0, _rice_cost(us, 0))
228
+ for k in range(1, MAX_K + 1):
229
+ cost = _rice_cost(us, k)
230
+ if cost < best[1]:
231
+ best = (k, cost)
232
+ return best
233
+
234
+
235
+ def _blocks(res: Sequence[int]) -> List[List[int]]:
236
+ us = [zigzag(r) for r in res]
237
+ return [us[i:i + BLOCK] for i in range(0, len(us), BLOCK)]
238
+
239
+
240
+ def _column_cost(ints: Sequence[int], order: int) -> int:
241
+ if len(ints) <= order:
242
+ return 1 << 60
243
+ cost = 3 # order field
244
+ for value in ints[:order]:
245
+ cost += 8 * max(1, (zigzag(value).bit_length() + 6) // 7)
246
+ for block in _blocks(residuals(ints, order)):
247
+ cost += K_BITS + _best_k(block)[1]
248
+ return cost
249
+
250
+
251
+ def choose_order(ints: Sequence[int]) -> int:
252
+ return min(range(MAX_ORDER + 1), key=lambda o: _column_cost(ints, o))
253
+
254
+
255
+ # ------------------------------------------------------------------ codec
256
+
257
+ MAGIC = b"PTC1"
258
+
259
+
260
+ def encode(csv_text: str):
261
+ lines = csv_text.rstrip("\n").split("\n")
262
+ header = lines[0].split(",")
263
+ rows = [line.split(",") for line in lines[1:]]
264
+ ncols, nrows = len(header), len(rows)
265
+
266
+ bw = BitWriter()
267
+ bw.uvarint(nrows)
268
+ bw.uvarint(ncols)
269
+ for name in header:
270
+ raw = name.encode("utf-8")
271
+ bw.uvarint(len(raw))
272
+ for byte in raw:
273
+ bw.bits(byte, 8)
274
+
275
+ report = []
276
+ for col in range(ncols):
277
+ cells = [row[col] for row in rows]
278
+ numeric = as_numeric_column(cells)
279
+ if numeric is None:
280
+ bw.bit(0)
281
+ for cell in cells:
282
+ raw = cell.encode("utf-8")
283
+ bw.uvarint(len(raw))
284
+ for byte in raw:
285
+ bw.bits(byte, 8)
286
+ report.append((header[col], None, None))
287
+ continue
288
+
289
+ ints, decimals = numeric
290
+ order = choose_order(ints)
291
+ bw.bit(1)
292
+ bw.bits(decimals, 6)
293
+ bw.bits(order, 3)
294
+ for value in ints[:order]:
295
+ bw.svarint(value)
296
+ for block in _blocks(residuals(ints, order)):
297
+ k, _ = _best_k(block)
298
+ bw.bits(k, K_BITS)
299
+ for u in block:
300
+ bw.rice(u, k)
301
+ report.append((header[col], order, _column_cost(ints, order) / 8.0))
302
+
303
+ return MAGIC + bw.bytes_out(), report
304
+
305
+
306
+ def decode(blob: bytes) -> str:
307
+ assert blob[:4] == MAGIC, "bad magic"
308
+ br = BitReader(blob[4:])
309
+ nrows = br.uvarint()
310
+ ncols = br.uvarint()
311
+ header = []
312
+ for _ in range(ncols):
313
+ n = br.uvarint()
314
+ header.append(bytes(br.bits(8) for _ in range(n)).decode("utf-8"))
315
+
316
+ columns = []
317
+ for _ in range(ncols):
318
+ if br.bit() == 0:
319
+ cells = []
320
+ for _ in range(nrows):
321
+ n = br.uvarint()
322
+ cells.append(bytes(br.bits(8) for _ in range(n)).decode("utf-8"))
323
+ columns.append(cells)
324
+ continue
325
+
326
+ decimals = br.bits(6)
327
+ order = br.bits(3)
328
+ warmup = [br.svarint() for _ in range(order)]
329
+ res = []
330
+ remaining = nrows - order
331
+ while remaining > 0:
332
+ k = br.bits(K_BITS)
333
+ take = min(BLOCK, remaining)
334
+ for _ in range(take):
335
+ res.append(unzigzag(br.rice(k)))
336
+ remaining -= take
337
+ ints = restore(warmup, res, order, nrows)
338
+ columns.append([int_to_cell(v, decimals) for v in ints])
339
+
340
+ out = [",".join(header)]
341
+ for i in range(nrows):
342
+ out.append(",".join(col[i] for col in columns))
343
+ return "\n".join(out) + "\n"