aboutsummaryrefslogtreecommitdiff
path: root/scripts/magic_finder.c
diff options
context:
space:
mode:
Diffstat (limited to 'scripts/magic_finder.c')
-rw-r--r--scripts/magic_finder.c224
1 files changed, 224 insertions, 0 deletions
diff --git a/scripts/magic_finder.c b/scripts/magic_finder.c
new file mode 100644
index 0000000..ee21587
--- /dev/null
+++ b/scripts/magic_finder.c
@@ -0,0 +1,224 @@
+#include <stdint.h>
+#include <stdio.h>
+#include <stdlib.h>
+#include <string.h>
+#include <time.h>
+
+// TODO: make magic bitboards use diffrent shift for each magic number (maybe)
+
+static int popcount(uint64_t b) { return __builtin_popcountll(b); }
+
+static uint64_t computeRookMask(int sq) {
+ uint64_t mask = 0;
+ int r = sq / 8, f = sq % 8;
+ for (int i = r + 1; i < 7; i++)
+ mask |= 1ULL << (i * 8 + f);
+ for (int i = r - 1; i > 0; i--)
+ mask |= 1ULL << (i * 8 + f);
+ for (int j = f + 1; j < 7; j++)
+ mask |= 1ULL << (r * 8 + j);
+ for (int j = f - 1; j > 0; j--)
+ mask |= 1ULL << (r * 8 + j);
+ return mask;
+}
+
+static uint64_t computeBishopMask(int sq) {
+ uint64_t mask = 0;
+ int r = sq / 8, f = sq % 8;
+ for (int i = r + 1, j = f + 1; i < 7 && j < 7; i++, j++)
+ mask |= 1ULL << (i * 8 + j);
+ for (int i = r + 1, j = f - 1; i < 7 && j > 0; i++, j--)
+ mask |= 1ULL << (i * 8 + j);
+ for (int i = r - 1, j = f + 1; i > 0 && j < 7; i--, j++)
+ mask |= 1ULL << (i * 8 + j);
+ for (int i = r - 1, j = f - 1; i > 0 && j > 0; i--, j--)
+ mask |= 1ULL << (i * 8 + j);
+ return mask;
+}
+
+static uint64_t computeRookAttacks(int sq, uint64_t occ) {
+ uint64_t att = 0;
+ int r = sq / 8, f = sq % 8;
+ for (int i = r + 1; i < 8; i++) {
+ att |= 1ULL << (i * 8 + f);
+ if (occ & (1ULL << (i * 8 + f)))
+ break;
+ }
+ for (int i = r - 1; i >= 0; i--) {
+ att |= 1ULL << (i * 8 + f);
+ if (occ & (1ULL << (i * 8 + f)))
+ break;
+ }
+ for (int j = f + 1; j < 8; j++) {
+ att |= 1ULL << (r * 8 + j);
+ if (occ & (1ULL << (r * 8 + j)))
+ break;
+ }
+ for (int j = f - 1; j >= 0; j--) {
+ att |= 1ULL << (r * 8 + j);
+ if (occ & (1ULL << (r * 8 + j)))
+ break;
+ }
+ return att;
+}
+
+static uint64_t computeBishopAttacks(int sq, uint64_t occ) {
+ uint64_t att = 0;
+ int r = sq / 8, f = sq % 8;
+ for (int i = r + 1, j = f + 1; i < 8 && j < 8; i++, j++) {
+ att |= 1ULL << (i * 8 + j);
+ if (occ & (1ULL << (i * 8 + j)))
+ break;
+ }
+ for (int i = r + 1, j = f - 1; i < 8 && j >= 0; i++, j--) {
+ att |= 1ULL << (i * 8 + j);
+ if (occ & (1ULL << (i * 8 + j)))
+ break;
+ }
+ for (int i = r - 1, j = f + 1; i >= 0 && j < 8; i--, j++) {
+ att |= 1ULL << (i * 8 + j);
+ if (occ & (1ULL << (i * 8 + j)))
+ break;
+ }
+ for (int i = r - 1, j = f - 1; i >= 0 && j >= 0; i--, j--) {
+ att |= 1ULL << (i * 8 + j);
+ if (occ & (1ULL << (i * 8 + j)))
+ break;
+ }
+ return att;
+}
+
+static uint64_t mapSquaresToMask(uint64_t index, uint64_t mask) {
+ uint64_t result = 0;
+ int bits = __builtin_popcountll(mask);
+ for (int i = 0; i < bits; i++) {
+ int sq = __builtin_ctzll(mask);
+ mask &= mask - 1;
+ if (index & (1ULL << i))
+ result |= (1ULL << sq);
+ }
+ return result;
+}
+
+static int findMagic(uint64_t mask, uint64_t (*attackFn)(int, uint64_t), int sq,
+ int shift, uint64_t *outMagic) {
+ int nbits = popcount(mask);
+ uint64_t nocc = 1ULL << nbits;
+
+ uint64_t *occs = malloc(nocc * sizeof(uint64_t));
+ uint64_t *atts = malloc(nocc * sizeof(uint64_t));
+ for (uint64_t i = 0; i < nocc; i++) {
+ occs[i] = mapSquaresToMask(i, mask);
+ atts[i] = attackFn(sq, occs[i]);
+ }
+
+ uint64_t tableSize = 1ULL << (64 - shift);
+ uint64_t *table = malloc(tableSize * sizeof(uint64_t));
+ uint64_t *epoch = calloc(tableSize, sizeof(uint64_t));
+ uint64_t cnt = 0;
+
+ for (unsigned long long attempt = 0; attempt < 10000000000ULL; attempt++) {
+ uint64_t magic;
+ do {
+ magic =
+ ((uint64_t)rand() << 32 | rand()) & ((uint64_t)rand() << 32 | rand());
+ } while (popcount((mask * magic) & 0xFF00000000000000ULL) < 6);
+
+ cnt++;
+ int ok = 1;
+ for (uint64_t i = 0; i < nocc; i++) {
+ uint64_t idx = (occs[i] * magic) >> shift;
+ if (epoch[idx] < cnt) {
+ epoch[idx] = cnt;
+ table[idx] = atts[i];
+ } else if (table[idx] != atts[i]) {
+ ok = 0;
+ break;
+ }
+ }
+ if (ok) {
+ *outMagic = magic;
+ free(occs);
+ free(atts);
+ free(table);
+ free(epoch);
+ return 1;
+ }
+ }
+
+ free(occs);
+ free(atts);
+ free(table);
+ free(epoch);
+ return 0;
+}
+
+int main(int argc, char *argv[]) {
+ int shift = 52;
+ if (argc > 1) {
+ shift = atoi(argv[1]);
+ if (shift < 32 || shift > 63) {
+ fprintf(stderr, "Usage: %s [shift] (32..63, default 52)\n", argv[0]);
+ return 1;
+ }
+ }
+ srand(time(NULL));
+
+ uint64_t rookMags[64], bishopMags[64];
+ uint64_t rookMasks[64], bishopMasks[64];
+
+ fprintf(stderr,
+ "Finding per-square rook magics (shift %d, table size %llu)...\n",
+ shift, (unsigned long long)(1ULL << (64 - shift)));
+ for (int sq = 0; sq < 64; sq++) {
+ rookMasks[sq] = computeRookMask(sq);
+ if (!findMagic(rookMasks[sq], computeRookAttacks, sq, shift,
+ &rookMags[sq])) {
+ fprintf(stderr, "FAILED to find rook magic for square %d\n", sq);
+ return 1;
+ }
+ fprintf(stderr, " sq %2d: %d bits, magic 0x%016llx\n", sq,
+ popcount(rookMasks[sq]), (unsigned long long)rookMags[sq]);
+ }
+
+ fprintf(stderr, "\nFinding per-square bishop magics (shift %d)...\n", shift);
+ for (int sq = 0; sq < 64; sq++) {
+ bishopMasks[sq] = computeBishopMask(sq);
+ if (!findMagic(bishopMasks[sq], computeBishopAttacks, sq, shift,
+ &bishopMags[sq])) {
+ fprintf(stderr, "FAILED to find bishop magic for square %d\n", sq);
+ return 1;
+ }
+ fprintf(stderr, " sq %2d: %d bits, magic 0x%016llx\n", sq,
+ popcount(bishopMasks[sq]), (unsigned long long)bishopMags[sq]);
+ }
+
+ fprintf(stderr, "\nDone! Writing magics.txt\n");
+
+ FILE *f = fopen("magics.txt", "w");
+ if (!f) {
+ fprintf(stderr, "Cannot open magics.txt\n");
+ return 1;
+ }
+
+ fprintf(f, "// shift = %d\n", shift);
+ fprintf(f, "constexpr int MAGIC_SHIFT = %d;\n\n", shift);
+
+ fprintf(f, "constexpr std::array<uint64_t, 64> rookMagics = {\n");
+ for (int i = 0; i < 64; i++) {
+ fprintf(f, " 0x%016llxULL%c // %c%d\n", (unsigned long long)rookMags[i],
+ i < 63 ? ',' : ' ', 'a' + i % 8, i / 8 + 1);
+ }
+ fprintf(f, "};\n\n");
+
+ fprintf(f, "constexpr std::array<uint64_t, 64> bishopMagics = {\n");
+ for (int i = 0; i < 64; i++) {
+ fprintf(f, " 0x%016llxULL%c // %c%d\n",
+ (unsigned long long)bishopMags[i], i < 63 ? ',' : ' ', 'a' + i % 8,
+ i / 8 + 1);
+ }
+ fprintf(f, "};\n");
+
+ fclose(f);
+ return 0;
+}