cm-bh-v2 / codec.py
dsaddsaf's picture
cm-bh-v2: context-mixing codec with bit-history state machines
fa5173a verified
Raw
History Blame Contribute Delete
25.7 kB
"""cm-bh-v2: context-mixing compressor, bit-history state machines (paq8l table).
Models (mixer inputs): exact order-1/2; hashed order-3/4/5/6/8 and two word
contexts via the 253-state paq8l bit-history FSM + 20-bit adaptive state maps
(count limit 1023); two sparse contexts; three match models (order-4/6/12);
six indirect (successor-history) contexts on orders 1/2/3/4/6 and words.
Mixing: 32 logistic weight sets (char class x match state x bit half) +
second bank keyed by high nibble, adaptive final blend.
SSE: paq8px-style series - p1=(p+3*APM(order0))/4, then blend of APMs over
order-1/2/3, match-state and word contexts.
All tables fixed-size (bounded memory regardless of input size).
Deterministic, CPU-only, pure Python stdlib.
"""
import math
from array import array
STRETCH_TABLE = [0.0] * 4097
for i in range(1, 4096):
p = i / 4096.0
STRETCH_TABLE[i] = math.log(p / (1 - p))
STRETCH_TABLE[0] = STRETCH_TABLE[1]
STRETCH_TABLE[4096] = STRETCH_TABLE[4095]
MATCH_ORDER = 4
MATCH_ORDER2 = 6
MATCH_ORDER3 = 12
HBITS = 25
HASH_MASK = (1 << HBITS) - 1
HASH_MASK6 = (1 << HBITS) - 1
WBITS = 22
WORD_MASK = (1 << WBITS) - 1
APM_CTX = 256
APM_CTX2 = 1024
APM_CTX3 = 64
APM_STEPS = 33
N_WSETS = 32
# ---- таблица состояний bit-history из paq8l/lpaq1 ----
N_STATES = 256
NEXT = bytearray(b'\x01\x02\x03\x05\x04\x06\x07\n\x08\x0c\t\r\x0b\x0e\x0f\x13\x10\x17\x11\x18\x12\x19\x14\x1b\x15\x1c\x16\x1d\x1a\x1e\x1f! # # # #"%"%"%"%"%"%$\'$\'$\'$\'&()+*-*-,/,/.1.10303245+69698;8;:=:=<?<?>A>A2BC7D9D9FIFIHKHKJMJMLOLO>Q>Q@RSETGTGVIVI,;,;:=:=<1<1LYLYN[N[P\\]E^W^W`-`-0c0cXeXePfgEhWhWj9j9>m>mXoXoPpqUrWrWt9t9>w>wXyXyZz{U|a|a~9~9>\x81>\x81b\x83b\x83Z\x84\x85U\x86a\x86a\x889\x889>\x8b>\x8bb\x8db\x8dZ\x8e\x8f_\x90a\x90aD9D9>Q>Qb\x93b\x93d\x94\x95_\x96k\x96kl\x97l\x97d\x98\x99_\x9akl\x9bd\x9c\x9d_\x9ekl\x9fd\xa0\xa1i\xa2kl\xa3n\xa4\xa5i\xa6uv\xa7n\xa8\xa9i\xaauv\xabn\xac\xadi\xaeuv\xafn\xb0\xb1i\xb2uv\xb3n\xb4\xb5s\xb6uv\xb7x\xb8\xb9s\xba\x7f\x80\xbbx\xbc\xbds\xbe\x7f\x80\xbfx\xc0\xc1s\xc2\x7f\x80\xc3x\xc4\xc5s\xc6\x7f\x80\xc7x\xc8\xc9s\xca\x7f\x80\xcbx\xcc\xcds\xce\x7f\x80\xcfx\xd0\xd1}\xd2\x7f\x80\xd3\x82\xd4\xd5}\xd6\x89\x8a\xd7\x82\xd8\xd9}\xda\x89\x8a\xdb\x82\xdc\xdd}\xde\x89\x8a\xdf\x82\xe0\xe1}\xe2\x89\x8a\xe3\x82\xe4\xe5}\xe6\x89\x8a\xe7\x82\xe8\xe9}\xea\x89\x8a\xeb\x82\xec\xed}\xee\x89\x8a\xef\x82\xf0\xf1}\xf2\x89\x8a\xf3\x82\xf4\xf5\x87\xf6\x89\x8a\xf7\x8c\xf8\xf9\x87\xfaEP\xfb\x8c\xfc\xf9\x87\xfaEP\xfb\x8c\xfc\x00\x00\x00\x00\x00\x00')
SM_INIT = [2048] * N_STATES
DSHIFT = bytes(2 if i < 7 else (3 if i < 31 else (4 if i < 120 else 5)) for i in range(N_STATES))
SM_LIMIT = 1023
SM_RCP = [int(65536 / (n + 1.5)) for n in range(SM_LIMIT + 1)] # скорость StateMap по счётчику
def make_sm():
return [p << 8 for p in SM_INIT], array('H', [0]) * N_STATES
RATE = bytes(2 if c < 2 else (3 if c < 6 else (4 if c < 30 else 5)) for c in range(256))
LR = 0.002
SHIFT = 4
APM_W1 = 3
APM_W2 = 3
APM_W3 = 2
MIX_W = 1
def squash(x):
if x > 20:
return 4095
if x < -20:
return 0
return int(4096 / (1 + math.exp(-x)))
class APM:
def __init__(self, n_ctx):
self.n = n_ctx
self.t = [0] * (n_ctx * APM_STEPS)
for c in range(n_ctx):
for i in range(APM_STEPS):
s = (i - (APM_STEPS // 2)) * (8.0 / (APM_STEPS // 2))
self.t[c * APM_STEPS + i] = squash(s) * 16
self.last_idx = 0
self.last_w = 0.0
def pp(self, p, ctx):
st = STRETCH_TABLE[p]
st = -8.0 if st < -8.0 else (8.0 if st > 8.0 else st)
pos = (st + 8.0) * ((APM_STEPS - 1) / 16.0)
lo = int(pos)
lo = APM_STEPS - 2 if lo >= APM_STEPS - 1 else lo
w = pos - lo
base = ctx * APM_STEPS + lo
v = self.t[base] * (1 - w) + self.t[base + 1] * w
self.last_idx = base
self.last_w = w
out = int(v / 16)
return 1 if out < 1 else (4094 if out > 4094 else out)
def update(self, bit):
g = bit * 4096 * 16
base = self.last_idx
self.t[base] += int((g - self.t[base]) * (1 - self.last_w) / 32)
self.t[base + 1] += int((g - self.t[base + 1]) * self.last_w / 32)
class Predictor:
def __init__(self, data_len_hint=0):
self.t1 = [2048] * (256 * 256)
self.t2 = [2048] * (1 << 24)
self.t3 = bytearray(HASH_MASK + 1)
self.t4 = bytearray(HASH_MASK + 1)
self.t6 = bytearray(HASH_MASK6 + 1)
self.succ1 = bytearray(256)
self.succ2 = bytearray(65536)
self.succ3 = bytearray(1 << 24)
self.succ4 = bytearray(1 << 24)
self.succ6 = bytearray(1 << 24)
self.succ1b = bytearray(256)
self.succw = bytearray(1 << 22)
self.ti1 = bytearray(1 << 22)
self.ti2 = bytearray(1 << 22)
self.ti3 = bytearray(1 << 22)
self.ti4 = bytearray(1 << 22)
self.ti6 = bytearray(1 << 22)
self.tiw = bytearray(1 << 22)
self.smi1 = make_sm()
self.smi2 = make_sm()
self.smi3 = make_sm()
self.smi4 = make_sm()
self.smi6 = make_sm()
self.smiw = make_sm()
self.sm3 = make_sm()
self.sm4 = make_sm()
self.sm6 = make_sm()
self.sm5 = make_sm()
self.sm8 = make_sm()
self.smw = make_sm()
self.smw1 = make_sm()
self.tw = bytearray(WORD_MASK + 1)
self.tw1 = bytearray(WORD_MASK + 1)
self.ts = [2048] * (256 * 256)
self.ts2 = [2048] * (256 * 256)
self.t5 = bytearray(WORD_MASK + 1)
self.t8 = bytearray(WORD_MASK + 1)
# weight sets: o1,o2,o3,o4,o6,match4,match8,word
base_w = [0.2, 0.2, 0.18, 0.18, 0.14, 0.3, 0.3, 0.2, 0.15, 0.1, 0.16, 0.35, 0.12, 0.12, 0.12, 0.08, 0.12, 0.12, 0.12, 0.12]
self.wsets = [list(base_w) for _ in range(N_WSETS)]
self.wsets_b = [list(base_w) for _ in range(16)]
self.wfinal = [0.5, 0.5]
self.c1 = 0
self.c2 = 0
self.c3 = 0
self.c4 = 0
self.c5 = 0
self.c8 = 0
self.c6 = 0
self.word_h = 0
self.prev_word_h = 0
self.history = bytearray(max(data_len_hint, 1))
self.pos = 0
self.match_table = [0] * (1 << 22)
self.match_ptr = -1
self.match_len = 0
self.match_table2 = [0] * (1 << 22)
self.match_ptr2 = -1
self.match_len2 = 0
self.match_table3 = [0] * (1 << 22)
self.match_ptr3 = -1
self.match_len3 = 0
self.apm0 = APM(256)
self.apm4 = APM(4096)
self.apm = APM(APM_CTX)
self.apm2 = APM(APM_CTX2)
self.apm3 = APM(APM_CTX3)
def _try_acquire_match(self):
if self.pos < MATCH_ORDER:
return None
h = 0
for i in range(MATCH_ORDER):
h = (h * 0x9E3779B1 + self.history[self.pos - MATCH_ORDER + i]) & 0xFFFFFFFF
h &= 0x3FFFFF
if self.match_ptr < 0:
v = self.match_table[h]
if v:
self.match_ptr = v - 1
self.match_len = 1
return h
def _try_acquire_match2(self):
if self.pos < MATCH_ORDER2:
return None
h = 0
for i in range(MATCH_ORDER2):
h = (h * 0x85EBCA6B + self.history[self.pos - MATCH_ORDER2 + i]) & 0xFFFFFFFF
h &= 0x3FFFFF
if self.match_ptr2 < 0:
v = self.match_table2[h]
if v:
self.match_ptr2 = v - 1
self.match_len2 = 1
return h
def _try_acquire_match3(self):
if self.pos < MATCH_ORDER3:
return None
h = 0
for i in range(MATCH_ORDER3):
h = (h * 0xC2B2AE3D + self.history[self.pos - MATCH_ORDER3 + i]) & 0xFFFFFFFF
h &= 0x3FFFFF
if self.match_ptr3 < 0:
v = self.match_table3[h]
if v:
self.match_ptr3 = v - 1
self.match_len3 = 1
return h
def step_byte(self, io_fn, byte_in=None):
match_hash = self._try_acquire_match()
match_hash2 = self._try_acquire_match2()
match_hash3 = self._try_acquire_match3()
predicted_byte = self.history[self.match_ptr] if self.match_ptr >= 0 else -1
predicted_byte2 = self.history[self.match_ptr2] if self.match_ptr2 >= 0 else -1
predicted_byte3 = self.history[self.match_ptr3] if self.match_ptr3 >= 0 else -1
node = 1
out_byte = 0
t1, t2, t3, t4, t6, tw = self.t1, self.t2, self.t3, self.t4, self.t6, self.tw
tw1, ts, t5, t8 = self.tw1, self.ts, self.t5, self.t8
ts2 = self.ts2
smp3, smn3 = self.sm3
smp4, smn4 = self.sm4
smp6, smn6 = self.sm6
smp5, smn5 = self.sm5
smp8, smn8 = self.sm8
smpw, smnw = self.smw
smpw1, smnw1 = self.smw1
NX = NEXT
RC = SM_RCP
ti1, ti2, ti3, ti4, ti6, tiw = self.ti1, self.ti2, self.ti3, self.ti4, self.ti6, self.tiw
smpi1, smni1 = self.smi1
smpi2, smni2 = self.smi2
smpi3, smni3 = self.smi3
smpi4, smni4 = self.smi4
smpi6, smni6 = self.smi6
smpiw, smniw = self.smiw
c1, c2, c3, c4, c5, c6, c8 = self.c1, self.c2, self.c3, self.c4, self.c5, self.c6, self.c8
match_len = self.match_len
match_len2 = self.match_len2
match_len3 = self.match_len3
apm0 = self.apm0
apm4 = self.apm4
apm = self.apm
apm2 = self.apm2
apm3 = self.apm3
# mixer weight-set selection: char class of previous byte x match states
if 97 <= c1 <= 122 or 65 <= c1 <= 90:
cls = 0
elif c1 == 32:
cls = 1
elif 48 <= c1 <= 57:
cls = 2
else:
cls = 3
msel = (1 if predicted_byte >= 0 else 0) + 2 * (1 if predicted_byte2 >= 0 else 0)
wlo = self.wsets[cls * 4 + msel]
whi = self.wsets[16 + cls * 4 + msel]
wb = self.wsets_b[c1 >> 4]
wf = self.wfinal
# word-model context for this byte
ci1 = (((c1 << 16) | (self.succ1[c1] << 8) | self.succ1b[c1]) * 0x9E3779B1) & 0xFFFFFFFFFFFF
ci2 = ((c2 << 8) | self.succ2[c2]) * 0x85EBCA6B
ci3 = ((c3 << 8) | self.succ3[c3]) * 0xC2B2AE3D
h4s = ((c4 * 0x2545F4914F6CDD1D) >> 24) & 0xFFFFFF
ci4 = ((c4 << 8) | self.succ4[h4s]) * 0x27D4EB2F
h6s = ((c6 * 0x9E3779B97F4A7C15) >> 28) & 0xFFFFFF
ci6 = (((c6 & 0xFFFFFFFF) << 8) | self.succ6[h6s]) * 0x165667B1
hws = (((self.word_h * 31) ^ self.prev_word_h) * 0x9E3779B1 >> 10) & 0x3FFFFF
ciw = ((self.word_h << 8) | self.succw[hws]) * 0x94D049BB
cw = ((self.word_h * 0x2545F4914F6CDD1D) ^ (self.prev_word_h * 0x9E3779B97F4A7C15)) & 0xFFFFFFFFFFFFFFFF
cw1 = (self.word_h * 0x9E3779B97F4A7C15) & 0xFFFFFFFFFFFFFFFF
sb = self.history[self.pos - 2] if self.pos >= 2 else 0
sb2 = self.history[self.pos - 3] if self.pos >= 3 else 0
for bitpos in range(8):
w = whi if bitpos < 4 else wlo
idx1 = (c1 << 8) | node
key2 = (c2 << 8) | node
idx3 = (((c3 * 0x2545F4914F6CDD1D) & 0xFFFFFFFFFFFFFFFF) ^ node) & HASH_MASK
idx4 = (((c4 * 0x2545F4914F6CDD1D) & 0xFFFFFFFFFFFFFFFF) ^ node) & HASH_MASK
idx6 = (((c6 * 0x9E3779B97F4A7C15) & 0xFFFFFFFFFFFFFFFF) ^ node) & HASH_MASK6
idx5 = (((c5 * 0xFF51AFD7ED558CCD) & 0xFFFFFFFFFFFFFFFF) ^ node) & WORD_MASK
idx8 = (((c8 * 0xD6E8FEB86659FD93) & 0xFFFFFFFFFFFFFFFF) ^ node) & WORD_MASK
idxw = (cw ^ node) & WORD_MASK
idxi1 = (ci1 ^ node) & 0x3FFFFF
idxi2 = (ci2 ^ node) & 0x3FFFFF
idxi3 = (ci3 ^ node) & 0x3FFFFF
idxi4 = (ci4 ^ node) & 0x3FFFFF
idxi6 = (ci6 ^ node) & 0x3FFFFF
idxiw = (ciw ^ node) & 0x3FFFFF
idxw1 = (cw1 ^ node) & WORD_MASK
idxs = (sb << 8) | node
idxs2 = (sb2 << 8) | node
p1 = t1[idx1]
p2 = t2[key2]
s3 = t3[idx3]
P3 = smp3[s3]
p3 = P3 >> 8
s4 = t4[idx4]
P4 = smp4[s4]
p4 = P4 >> 8
s6 = t6[idx6]
P6 = smp6[s6]
p6 = P6 >> 8
s5 = t5[idx5]
P5 = smp5[s5]
p5 = P5 >> 8
s8 = t8[idx8]
P8 = smp8[s8]
p8 = P8 >> 8
sw = tw[idxw]
Pw = smpw[sw]
pw = Pw >> 8
si1 = ti1[idxi1]
Pi1 = smpi1[si1]
pi1 = Pi1 >> 8
si2 = ti2[idxi2]
Pi2 = smpi2[si2]
pi2 = Pi2 >> 8
si3 = ti3[idxi3]
Pi3 = smpi3[si3]
pi3 = Pi3 >> 8
si4 = ti4[idxi4]
Pi4 = smpi4[si4]
pi4 = Pi4 >> 8
si6 = ti6[idxi6]
Pi6i = smpi6[si6]
pi6i = Pi6i >> 8
siw = tiw[idxiw]
Piw = smpiw[siw]
piw = Piw >> 8
sw1 = tw1[idxw1]
Pw1 = smpw1[sw1]
pw1 = Pw1 >> 8
psp = ts[idxs]
psp2 = ts2[idxs2]
st1 = STRETCH_TABLE[p1]
st2 = STRETCH_TABLE[p2]
st3 = STRETCH_TABLE[p3]
st4 = STRETCH_TABLE[p4]
st6 = STRETCH_TABLE[p6]
st5 = STRETCH_TABLE[p5]
st8 = STRETCH_TABLE[p8]
stw = STRETCH_TABLE[pw]
sti1 = STRETCH_TABLE[pi1]
sti2 = STRETCH_TABLE[pi2]
sti3 = STRETCH_TABLE[pi3]
sti4 = STRETCH_TABLE[pi4]
sti6 = STRETCH_TABLE[pi6i]
stiw = STRETCH_TABLE[piw]
stw1 = STRETCH_TABLE[pw1]
stsp = STRETCH_TABLE[psp]
stsp2 = STRETCH_TABLE[psp2]
if predicted_byte >= 0:
pred_bit = (predicted_byte >> (7 - bitpos)) & 1
conf = min(3.5, 0.3 + match_len * 0.25)
stm = conf if pred_bit == 1 else -conf
else:
stm = 0.0
if predicted_byte2 >= 0:
pred_bit2 = (predicted_byte2 >> (7 - bitpos)) & 1
conf2 = min(5.0, 0.5 + match_len2 * 0.35)
stm2 = conf2 if pred_bit2 == 1 else -conf2
else:
stm2 = 0.0
if predicted_byte3 >= 0:
pred_bit3 = (predicted_byte3 >> (7 - bitpos)) & 1
conf3 = min(6.5, 1.0 + match_len3 * 0.5)
stm3 = conf3 if pred_bit3 == 1 else -conf3
else:
stm3 = 0.0
dot = (w[0] * st1 + w[1] * st2 + w[2] * st3 + w[3] * st4 + w[4] * st6
+ w[5] * stm + w[6] * stm2 + w[7] * stw + w[8] * stw1 + w[9] * stsp + w[10] * st5 + w[11] * stm3 + w[12] * st8 + w[13] * sti1 + w[14] * sti2 + w[15] * stsp2 + w[16] * sti3 + w[17] * sti4 + w[18] * sti6 + w[19] * stiw)
dot_b = (wb[0] * st1 + wb[1] * st2 + wb[2] * st3 + wb[3] * st4 + wb[4] * st6
+ wb[5] * stm + wb[6] * stm2 + wb[7] * stw + wb[8] * stw1 + wb[9] * stsp + wb[10] * st5 + wb[11] * stm3 + wb[12] * st8 + wb[13] * sti1 + wb[14] * sti2 + wb[15] * stsp2 + wb[16] * sti3 + wb[17] * sti4 + wb[18] * sti6 + wb[19] * stiw)
fdot = wf[0] * dot + wf[1] * dot_b
mixed = squash(fdot)
mixed = 1 if mixed < 1 else (4094 if mixed > 4094 else mixed)
m1 = (mixed + 3 * apm0.pp(mixed, node)) >> 2
m1 = 1 if m1 < 1 else (4094 if m1 > 4094 else m1)
r1 = apm.pp(m1, c1)
r2 = apm2.pp(m1, c2 & (APM_CTX2 - 1))
mstate = (0 if match_len2 == 0 else (1 if match_len2 < 8 else (2 if match_len2 < 24 else 3)))
r3 = apm3.pp(m1, (mstate << 2) | cls)
r4 = apm4.pp(m1, ((c3 * 0x9E3779B1) >> 20) & 4095)
final_p = (r1 + 2 * r2 + r3 + 2 * r4) // 6
final_p = 1 if final_p < 1 else (4094 if final_p > 4094 else final_p)
bit = io_fn(final_p, byte_in, bitpos)
err = (bit * 4096 - mixed) / 4096.0
lr = LR
w[0] += lr * err * st1
w[1] += lr * err * st2
w[2] += lr * err * st3
w[3] += lr * err * st4
w[4] += lr * err * st6
w[5] += lr * err * stm
w[6] += lr * err * stm2
w[7] += lr * err * stw
w[8] += lr * err * stw1
w[9] += lr * err * stsp
w[10] += lr * err * st5
w[11] += lr * err * stm3
w[12] += lr * err * st8
wb[0] += lr * err * st1
wb[1] += lr * err * st2
wb[2] += lr * err * st3
wb[3] += lr * err * st4
wb[4] += lr * err * st6
wb[5] += lr * err * stm
wb[6] += lr * err * stm2
wb[7] += lr * err * stw
wb[8] += lr * err * stw1
wb[9] += lr * err * stsp
wb[10] += lr * err * st5
wb[11] += lr * err * stm3
wb[12] += lr * err * st8
w[13] += lr * err * sti1
w[14] += lr * err * sti2
wb[13] += lr * err * sti1
wb[14] += lr * err * sti2
w[15] += lr * err * stsp2
w[16] += lr * err * sti3
wb[15] += lr * err * stsp2
wb[16] += lr * err * sti3
w[17] += lr * err * sti4
wb[17] += lr * err * sti4
w[18] += lr * err * sti6
wb[18] += lr * err * sti6
w[19] += lr * err * stiw
wb[19] += lr * err * stiw
wf[0] += 0.0005 * err * dot
wf[1] += 0.0005 * err * dot_b
apm0.update(bit)
apm4.update(bit)
apm.update(bit)
apm2.update(bit)
apm3.update(bit)
tgt = bit * 4096
tgth = bit * 1048320
t1[idx1] = p1 + ((tgt - p1) >> SHIFT)
t2[key2] = p2 + ((tgt - p2) >> SHIFT)
c = smn3[s3]
smp3[s3] = P3 + (((tgth - P3) * RC[c]) >> 16)
if c < SM_LIMIT: smn3[s3] = c + 1
t3[idx3] = NX[(s3 << 1) | bit]
c = smn4[s4]
smp4[s4] = P4 + (((tgth - P4) * RC[c]) >> 16)
if c < SM_LIMIT: smn4[s4] = c + 1
t4[idx4] = NX[(s4 << 1) | bit]
c = smn6[s6]
smp6[s6] = P6 + (((tgth - P6) * RC[c]) >> 16)
if c < SM_LIMIT: smn6[s6] = c + 1
t6[idx6] = NX[(s6 << 1) | bit]
c = smn5[s5]
smp5[s5] = P5 + (((tgth - P5) * RC[c]) >> 16)
if c < SM_LIMIT: smn5[s5] = c + 1
t5[idx5] = NX[(s5 << 1) | bit]
c = smn8[s8]
smp8[s8] = P8 + (((tgth - P8) * RC[c]) >> 16)
if c < SM_LIMIT: smn8[s8] = c + 1
t8[idx8] = NX[(s8 << 1) | bit]
c = smni1[si1]
smpi1[si1] = Pi1 + (((tgth - Pi1) * RC[c]) >> 16)
if c < SM_LIMIT: smni1[si1] = c + 1
ti1[idxi1] = NX[(si1 << 1) | bit]
c = smni2[si2]
smpi2[si2] = Pi2 + (((tgth - Pi2) * RC[c]) >> 16)
if c < SM_LIMIT: smni2[si2] = c + 1
ti2[idxi2] = NX[(si2 << 1) | bit]
c = smni3[si3]
smpi3[si3] = Pi3 + (((tgth - Pi3) * RC[c]) >> 16)
if c < SM_LIMIT: smni3[si3] = c + 1
ti3[idxi3] = NX[(si3 << 1) | bit]
c = smni4[si4]
smpi4[si4] = Pi4 + (((tgth - Pi4) * RC[c]) >> 16)
if c < SM_LIMIT: smni4[si4] = c + 1
ti4[idxi4] = NX[(si4 << 1) | bit]
c = smni6[si6]
smpi6[si6] = Pi6i + (((tgth - Pi6i) * RC[c]) >> 16)
if c < SM_LIMIT: smni6[si6] = c + 1
ti6[idxi6] = NX[(si6 << 1) | bit]
c = smniw[siw]
smpiw[siw] = Piw + (((tgth - Piw) * RC[c]) >> 16)
if c < SM_LIMIT: smniw[siw] = c + 1
tiw[idxiw] = NX[(siw << 1) | bit]
c = smnw[sw]
smpw[sw] = Pw + (((tgth - Pw) * RC[c]) >> 16)
if c < SM_LIMIT: smnw[sw] = c + 1
tw[idxw] = NX[(sw << 1) | bit]
c = smnw1[sw1]
smpw1[sw1] = Pw1 + (((tgth - Pw1) * RC[c]) >> 16)
if c < SM_LIMIT: smnw1[sw1] = c + 1
tw1[idxw1] = NX[(sw1 << 1) | bit]
ts[idxs] = psp + ((tgt - psp) >> SHIFT)
ts2[idxs2] = psp2 + ((tgt - psp2) >> SHIFT)
if predicted_byte >= 0 and bit != pred_bit:
predicted_byte = -1
if predicted_byte2 >= 0 and bit != pred_bit2:
predicted_byte2 = -1
if predicted_byte3 >= 0 and bit != pred_bit3:
predicted_byte3 = -1
node = (node << 1) | bit
out_byte = (out_byte << 1) | bit
if self.pos >= len(self.history):
self.history.append(out_byte)
else:
self.history[self.pos] = out_byte
if self.match_ptr >= 0 and self.match_ptr < self.pos and self.history[self.match_ptr] == out_byte:
self.match_ptr += 1
self.match_len += 1
else:
self.match_ptr = -1
self.match_len = 0
if match_hash is not None:
self.match_table[match_hash] = self.pos + 1
if self.match_ptr2 >= 0 and self.match_ptr2 < self.pos and self.history[self.match_ptr2] == out_byte:
self.match_ptr2 += 1
self.match_len2 += 1
else:
self.match_ptr2 = -1
self.match_len2 = 0
if match_hash2 is not None:
self.match_table2[match_hash2] = self.pos + 1
if self.match_ptr3 >= 0 and self.match_ptr3 < self.pos and self.history[self.match_ptr3] == out_byte:
self.match_ptr3 += 1
self.match_len3 += 1
else:
self.match_ptr3 = -1
self.match_len3 = 0
if match_hash3 is not None:
self.match_table3[match_hash3] = self.pos + 1
# word model state
b = out_byte
if 97 <= b <= 122 or 65 <= b <= 90:
self.word_h = (self.word_h * 0x9E3779B1 + (b | 0x20)) & 0xFFFFFFFFFFFFFFFF
elif self.word_h:
self.prev_word_h = self.word_h
self.word_h = 0
self.succ1b[c1] = self.succ1[c1]
self.succ1[c1] = out_byte
self.succ2[c2] = out_byte
self.succ3[c3] = out_byte
self.succ4[h4s] = out_byte
self.succ6[h6s] = out_byte
self.succw[hws] = out_byte
self.c2 = ((c1 << 8) | out_byte) & 0xFFFF
self.c3 = ((c3 << 8) | out_byte) & 0xFFFFFF
self.c4 = ((c4 << 8) | out_byte) & 0xFFFFFFFF
self.c5 = ((c5 << 8) | out_byte) & 0xFFFFFFFFFF
self.c8 = ((c8 << 8) | out_byte) & 0xFFFFFFFFFFFFFFFF
self.c6 = ((c6 << 8) | out_byte) & 0xFFFFFFFFFFFF
self.c1 = out_byte
self.pos += 1
return out_byte
TOP = 1 << 24
BOT = 1 << 16
MASK32 = (1 << 32) - 1
class RangeEncoderBit:
def __init__(self):
self.low = 0
self.range = MASK32
self.out = bytearray()
def encode_bit(self, p1, bit):
mid = (self.range >> 12) * p1
if bit:
self.range = mid
else:
self.low = (self.low + mid) & MASK32
self.range -= mid
while True:
if (self.low ^ (self.low + self.range)) < TOP:
pass
elif self.range < BOT:
self.range = (-self.low) & (BOT - 1)
else:
break
self.out.append((self.low >> 24) & 0xFF)
self.low = (self.low << 8) & MASK32
self.range = (self.range << 8) & MASK32
def finish(self):
for _ in range(4):
self.out.append((self.low >> 24) & 0xFF)
self.low = (self.low << 8) & MASK32
return bytes(self.out)
class RangeDecoderBit:
def __init__(self, data):
self.data = data
self.pos = 0
self.low = 0
self.range = MASK32
self.code = 0
for _ in range(4):
self.code = ((self.code << 8) | self._byte()) & MASK32
def _byte(self):
if self.pos < len(self.data):
b = self.data[self.pos]
self.pos += 1
return b
return 0
def decode_bit(self, p1):
mid = (self.range >> 12) * p1
if (self.code - self.low) & MASK32 < mid:
bit = 1
self.range = mid
else:
bit = 0
self.low = (self.low + mid) & MASK32
self.range -= mid
while True:
if (self.low ^ (self.low + self.range)) < TOP:
pass
elif self.range < BOT:
self.range = (-self.low) & (BOT - 1)
else:
break
self.code = ((self.code << 8) | self._byte()) & MASK32
self.low = (self.low << 8) & MASK32
self.range = (self.range << 8) & MASK32
return bit
def compress(data: bytes) -> bytes:
import gc
gc.disable()
pred = Predictor(len(data))
enc = RangeEncoderBit()
def io_fn(p, byte_in, bitpos):
bit = (byte_in >> (7 - bitpos)) & 1
enc.encode_bit(p, bit)
return bit
for b in data:
pred.step_byte(io_fn, b)
body = enc.finish()
return len(data).to_bytes(8, "big") + body
def decompress(blob: bytes) -> bytes:
import gc
gc.disable()
n = int.from_bytes(blob[:8], "big")
dec = RangeDecoderBit(blob[8:])
pred = Predictor(n)
out = bytearray()
def io_fn(p, byte_in, bitpos):
return dec.decode_bit(p)
for _ in range(n):
b = pred.step_byte(io_fn)
out.append(b)
return bytes(out)