diff options
Diffstat (limited to 'scripts/magic_finder.c')
| -rw-r--r-- | scripts/magic_finder.c | 224 |
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; +} |
