Add decay to ANS predictor and further optimize it
This commit is contained in:
+36
-36
@@ -182,7 +182,7 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": 123,
|
"execution_count": 4,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [
|
"outputs": [
|
||||||
{
|
{
|
||||||
@@ -191,7 +191,7 @@
|
|||||||
"'-0.91, -0.90, -0.89, -0.88, -0.86, -0.83, -0.80, -0.75, -0.67, -0.50, 0.00, 0.50, 0.67, 0.75, 0.80, 0.83, 0.86, 0.88, 0.89, 0.90, 0.91'"
|
"'-0.91, -0.90, -0.89, -0.88, -0.86, -0.83, -0.80, -0.75, -0.67, -0.50, 0.00, 0.50, 0.67, 0.75, 0.80, 0.83, 0.86, 0.88, 0.89, 0.90, 0.91'"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
"execution_count": 123,
|
"execution_count": 4,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"output_type": "execute_result"
|
"output_type": "execute_result"
|
||||||
}
|
}
|
||||||
@@ -245,35 +245,35 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": 128,
|
"execution_count": 6,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [
|
"outputs": [
|
||||||
{
|
{
|
||||||
"name": "stdout",
|
"name": "stdout",
|
||||||
"output_type": "stream",
|
"output_type": "stream",
|
||||||
"text": [
|
"text": [
|
||||||
"p=PosixPath('../corpus/quantized/floodedcaves_0.png'), bytes=2315 (bits=18520, bits per pixel=1.13037109375)\n",
|
"p=PosixPath('../corpus/quantized/floodedcaves_0.png'), bytes=1743 (percent=21.3%, bits per pixel=0.85107421875)\n",
|
||||||
"p=PosixPath('../corpus/quantized/age_of_ants-title.png'), bytes=4301 (bits=34408, bits per pixel=2.10009765625)\n",
|
"p=PosixPath('../corpus/quantized/age_of_ants-title.png'), bytes=3222 (percent=39.3%, bits per pixel=1.5732421875)\n",
|
||||||
"p=PosixPath('../corpus/quantized/build a jetpack_0.png'), bytes=1804 (bits=14432, bits per pixel=0.880859375)\n",
|
"p=PosixPath('../corpus/quantized/build a jetpack_0.png'), bytes=1223 (percent=14.9%, bits per pixel=0.59716796875)\n",
|
||||||
"p=PosixPath('../corpus/quantized/shinkansen 1 3_1.png'), bytes=1977 (bits=15816, bits per pixel=0.96533203125)\n",
|
"p=PosixPath('../corpus/quantized/shinkansen 1 3_1.png'), bytes=1178 (percent=14.4%, bits per pixel=0.5751953125)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_build_a_jetpack-1.png'), bytes=1722 (bits=13776, bits per pixel=0.8408203125)\n",
|
"p=PosixPath('../corpus/quantized/pico8_build_a_jetpack-1.png'), bytes=1116 (percent=13.6%, bits per pixel=0.544921875)\n",
|
||||||
"p=PosixPath('../corpus/quantized/iso_ray-09_000.png'), bytes=4857 (bits=38856, bits per pixel=2.37158203125)\n",
|
"p=PosixPath('../corpus/quantized/iso_ray-09_000.png'), bytes=3614 (percent=44.1%, bits per pixel=1.7646484375)\n",
|
||||||
"p=PosixPath('../corpus/quantized/age of ants_2.png'), bytes=3170 (bits=25360, bits per pixel=1.5478515625)\n",
|
"p=PosixPath('../corpus/quantized/age of ants_2.png'), bytes=2190 (percent=26.7%, bits per pixel=1.0693359375)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_mot_raymarch3-0.png'), bytes=5249 (bits=41992, bits per pixel=2.56298828125)\n",
|
"p=PosixPath('../corpus/quantized/pico8_mot_raymarch3-0.png'), bytes=4233 (percent=51.7%, bits per pixel=2.06689453125)\n",
|
||||||
"p=PosixPath('../corpus/quantized/age of ants_5.png'), bytes=2831 (bits=22648, bits per pixel=1.38232421875)\n",
|
"p=PosixPath('../corpus/quantized/age of ants_5.png'), bytes=1988 (percent=24.3%, bits per pixel=0.970703125)\n",
|
||||||
"p=PosixPath('../corpus/quantized/donsol for pico-8 v1 8 4 _2.png'), bytes=1684 (bits=13472, bits per pixel=0.822265625)\n",
|
"p=PosixPath('../corpus/quantized/donsol for pico-8 v1 8 4 _2.png'), bytes=1006 (percent=12.3%, bits per pixel=0.4912109375)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_witchcrafttd-6.png'), bytes=1759 (bits=14072, bits per pixel=0.85888671875)\n",
|
"p=PosixPath('../corpus/quantized/pico8_witchcrafttd-6.png'), bytes=1272 (percent=15.5%, bits per pixel=0.62109375)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_bunnysurvivor-9.png'), bytes=1738 (bits=13904, bits per pixel=0.8486328125)\n",
|
"p=PosixPath('../corpus/quantized/pico8_bunnysurvivor-9.png'), bytes=1189 (percent=14.5%, bits per pixel=0.58056640625)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_rotslimepires_1_1-0.png'), bytes=5030 (bits=40240, bits per pixel=2.4560546875)\n",
|
"p=PosixPath('../corpus/quantized/pico8_rotslimepires_1_1-0.png'), bytes=3144 (percent=38.4%, bits per pixel=1.53515625)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_donsol8_v1-14.png'), bytes=2474 (bits=19792, bits per pixel=1.2080078125)\n",
|
"p=PosixPath('../corpus/quantized/pico8_donsol8_v1-14.png'), bytes=1236 (percent=15.1%, bits per pixel=0.603515625)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_ppwr-5.png'), bytes=3169 (bits=25352, bits per pixel=1.54736328125)\n",
|
"p=PosixPath('../corpus/quantized/pico8_ppwr-5.png'), bytes=1742 (percent=21.3%, bits per pixel=0.8505859375)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico8_px9-9.png'), bytes=1699 (bits=13592, bits per pixel=0.82958984375)\n",
|
"p=PosixPath('../corpus/quantized/pico8_px9-9.png'), bytes=1064 (percent=13.0%, bits per pixel=0.51953125)\n",
|
||||||
"p=PosixPath('../corpus/quantized/build a jetpack_5.png'), bytes=1677 (bits=13416, bits per pixel=0.81884765625)\n",
|
"p=PosixPath('../corpus/quantized/build a jetpack_5.png'), bytes=1117 (percent=13.6%, bits per pixel=0.54541015625)\n",
|
||||||
"p=PosixPath('../corpus/quantized/pico20068.png'), bytes=2601 (bits=20808, bits per pixel=1.27001953125)\n",
|
"p=PosixPath('../corpus/quantized/pico20068.png'), bytes=1464 (percent=17.9%, bits per pixel=0.71484375)\n",
|
||||||
"p=PosixPath('../corpus/quantized/storming the grandmothership_1.png'), bytes=3602 (bits=28816, bits per pixel=1.7587890625)\n",
|
"p=PosixPath('../corpus/quantized/storming the grandmothership_1.png'), bytes=2231 (percent=27.2%, bits per pixel=1.08935546875)\n",
|
||||||
"p=PosixPath('../corpus/quantized/build a jetpack_9.png'), bytes=1773 (bits=14184, bits per pixel=0.86572265625)\n",
|
"p=PosixPath('../corpus/quantized/build a jetpack_9.png'), bytes=1315 (percent=16.1%, bits per pixel=0.64208984375)\n",
|
||||||
"p=PosixPath('../corpus/quantized/hersheys_train_line_0.png'), bytes=4459 (bits=35672, bits per pixel=2.17724609375)\n",
|
"p=PosixPath('../corpus/quantized/hersheys_train_line_0.png'), bytes=2613 (percent=31.9%, bits per pixel=1.27587890625)\n",
|
||||||
"Total bytes: 59891\n"
|
"Total bytes: 39900\n"
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
@@ -281,9 +281,10 @@
|
|||||||
"class Predictor:\n",
|
"class Predictor:\n",
|
||||||
" def __init__(self):\n",
|
" def __init__(self):\n",
|
||||||
" self.counts = [Counter() for _ in range(4)]\n",
|
" self.counts = [Counter() for _ in range(4)]\n",
|
||||||
" gain = 0.020\n",
|
" gain = 1.27\n",
|
||||||
" weights = [2, 100, 2, 100]\n",
|
" weights = [30, 100] * 2\n",
|
||||||
" self.weights = [x/sum(weights) * gain for x in weights]\n",
|
" self.weights = [x/sum(weights) * gain for x in weights]\n",
|
||||||
|
" self.decay = 0.9\n",
|
||||||
"\n",
|
"\n",
|
||||||
" def _keys(self, neighbors, partial_pixel):\n",
|
" def _keys(self, neighbors, partial_pixel):\n",
|
||||||
" assert len(neighbors) == 4\n",
|
" assert len(neighbors) == 4\n",
|
||||||
@@ -292,8 +293,6 @@
|
|||||||
" def predict(self, contexts):\n",
|
" def predict(self, contexts):\n",
|
||||||
" \"\"\"Returns probability that the next bit is zero, out of 256\"\"\"\n",
|
" \"\"\"Returns probability that the next bit is zero, out of 256\"\"\"\n",
|
||||||
" total = 0.0\n",
|
" total = 0.0\n",
|
||||||
" # for i in range(4):\n",
|
|
||||||
" # total += self.weights[i] * self.counts[i][contexts[i]]\n",
|
|
||||||
" for counter, key, weight in zip(self.counts, contexts, self.weights):\n",
|
" for counter, key, weight in zip(self.counts, contexts, self.weights):\n",
|
||||||
" total += weight * counter[key]\n",
|
" total += weight * counter[key]\n",
|
||||||
" prob_float = (sigmoid(total) + 1.0) / 2.0\n",
|
" prob_float = (sigmoid(total) + 1.0) / 2.0\n",
|
||||||
@@ -305,7 +304,7 @@
|
|||||||
" # Note for the future: code might be simpler if we used P(bit=1) everywhere\n",
|
" # Note for the future: code might be simpler if we used P(bit=1) everywhere\n",
|
||||||
" delta = -int(bit) * 2 + 1\n",
|
" delta = -int(bit) * 2 + 1\n",
|
||||||
" for counter, key in zip(self.counts, contexts):\n",
|
" for counter, key in zip(self.counts, contexts):\n",
|
||||||
" counter[key] += delta\n",
|
" counter[key] = counter[key] * self.decay + delta\n",
|
||||||
"\n",
|
"\n",
|
||||||
"\n",
|
"\n",
|
||||||
"def bitwise_encode(img: np.array):\n",
|
"def bitwise_encode(img: np.array):\n",
|
||||||
@@ -352,7 +351,7 @@
|
|||||||
" img = np.array(Image.open(p))\n",
|
" img = np.array(Image.open(p))\n",
|
||||||
" t = len(bitwise_encode(img))\n",
|
" t = len(bitwise_encode(img))\n",
|
||||||
" byte_size += t\n",
|
" byte_size += t\n",
|
||||||
" print(f\"{p=}, bytes={t} (bits={t*8}, bits per pixel={t*8/len(img.flatten())})\")\n",
|
" print(f\"{p=}, bytes={t} (percent={t/8192*100:.3}%, bits per pixel={t*8/128**2})\")\n",
|
||||||
" print(\"Total bytes:\", byte_size)\n",
|
" print(\"Total bytes:\", byte_size)\n",
|
||||||
"\n",
|
"\n",
|
||||||
"corpus_bitwise_encode()\n"
|
"corpus_bitwise_encode()\n"
|
||||||
@@ -365,12 +364,13 @@
|
|||||||
"Things learned so far:\n",
|
"Things learned so far:\n",
|
||||||
"\n",
|
"\n",
|
||||||
"- ANS is promising!\n",
|
"- ANS is promising!\n",
|
||||||
"- Diagonal neighbors aren't super useful as context, compared to orthogonal neighbors\n",
|
"- Diagonal neighbors currently count for 30% relative to orthogonal neighbors.\n",
|
||||||
"- Compression bit-by-bit (no RLE, etc) gets us 34.8% compression: definitely doing something, but not great on its own\n",
|
"- It's important to bound counts (e.g., with a decay) so they don't go out of control.\n",
|
||||||
"- From eyeballing failed predictions, we might benefit from clamping the probability values\n",
|
" - It looks like we had been confidently wrong (1-15 or 240-255) about 3-4% of the time?\n",
|
||||||
" - It looks like we're confidently wrong (1-15 or 240-255) about 3-4% of the time?\n",
|
|
||||||
" - Assuming a wrong guess costs 6 bits and 2500 wrong guesses/image: this costs ~2k per image or ~40k for the corpus\n",
|
" - Assuming a wrong guess costs 6 bits and 2500 wrong guesses/image: this costs ~2k per image or ~40k for the corpus\n",
|
||||||
" - If we don't need the precision, we could also just reduce probability granularities"
|
" - We shaved off about half that (~20k) by adding decay and fine-tuning it\n",
|
||||||
|
"- If we don't need the precision, we could reduce probability granularities.\n",
|
||||||
|
"- Compression bit-by-bit (no RLE, etc) gets us 23.2% compression on average. This sometimes beats PX9!"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user