Add decay to ANS predictor and further optimize it

This commit is contained in:
2024-12-02 00:22:40 -08:00
parent d789388f55
commit a039b54c41
+36 -36
View File
@@ -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!"
] ]
}, },
{ {