#include "engine/magic.h" #include "random.h" #include #include #include static magic_t rook_table[64]; static magic_t bishop_table[64]; static bitboard_t *rook_attacks_flat; static bitboard_t *bishop_attacks_flat; static bitboard_t compute_rook_mask(int sq) { bitboard_t mask = 0; int rank = sq / 8, file = sq % 8; for (int r = rank + 1; r < 7; r++) mask |= (bitboard_t)1 << (r * 8 + file); for (int r = rank - 1; r > 0; r--) mask |= (bitboard_t)1 << (r * 8 + file); for (int f = file + 1; f < 7; f++) mask |= (bitboard_t)1 << (rank * 8 + f); for (int f = file - 1; f > 0; f--) mask |= (bitboard_t)1 << (rank * 8 + f); return mask; } static bitboard_t compute_bishop_mask(int sq) { bitboard_t mask = 0; int rank = sq / 8, file = sq % 8; for (int r = rank + 1, f = file + 1; r < 7 && f < 7; r++, f++) mask |= (bitboard_t)1 << (r * 8 + f); for (int r = rank - 1, f = file - 1; r > 0 && f > 0; r--, f--) mask |= (bitboard_t)1 << (r * 8 + f); for (int r = rank - 1, f = file + 1; r > 0 && f < 7; r--, f++) mask |= (bitboard_t)1 << (r * 8 + f); for (int r = rank + 1, f = file - 1; r < 7 && f > 0; r++, f--) mask |= (bitboard_t)1 << (r * 8 + f); return mask; } static bitboard_t compute_rook_attacks_ref(int sq, bitboard_t occ) { bitboard_t attacks = 0; int rank = sq / 8, file = sq % 8; for (int r = rank + 1; r < 8; r++) { bitboard_t b = (bitboard_t)1 << (r * 8 + file); attacks |= b; if (occ & b) break; } for (int r = rank - 1; r >= 0; r--) { bitboard_t b = (bitboard_t)1 << (r * 8 + file); attacks |= b; if (occ & b) break; } for (int f = file + 1; f < 8; f++) { bitboard_t b = (bitboard_t)1 << (rank * 8 + f); attacks |= b; if (occ & b) break; } for (int f = file - 1; f >= 0; f--) { bitboard_t b = (bitboard_t)1 << (rank * 8 + f); attacks |= b; if (occ & b) break; } return attacks; } static bitboard_t compute_bishop_attacks_ref(int sq, bitboard_t occ) { bitboard_t attacks = 0; int rank = sq / 8, file = sq % 8; for (int r = rank + 1, f = file + 1; r < 8 && f < 8; r++, f++) { bitboard_t b = (bitboard_t)1 << (r * 8 + f); attacks |= b; if (occ & b) break; } for (int r = rank - 1, f = file - 1; r >= 0 && f >= 0; r--, f--) { bitboard_t b = (bitboard_t)1 << (r * 8 + f); attacks |= b; if (occ & b) break; } for (int r = rank - 1, f = file + 1; r >= 0 && f < 8; r--, f++) { bitboard_t b = (bitboard_t)1 << (r * 8 + f); attacks |= b; if (occ & b) break; } for (int r = rank + 1, f = file - 1; r < 8 && f >= 0; r++, f--) { bitboard_t b = (bitboard_t)1 << (r * 8 + f); attacks |= b; if (occ & b) break; } return attacks; } static void init_magic_for_square( int sq, magic_t *entry, bool is_bishop, bitboard_t *occs, bitboard_t *atts, bitboard_t *used ) { int bits = 64 - entry->shift; int count = 1 << bits; bitboard_t mask = entry->mask; for (int i = 0; i < count; i++) { occs[i] = 0; bitboard_t m = mask; for (int j = 0; m; j++) { if (i & (1 << j)) { occs[i] |= (bitboard_t)1 << __builtin_ctzll(m); } m &= m - 1; } atts[i] = is_bishop ? compute_bishop_attacks_ref(sq, occs[i]) : compute_rook_attacks_ref(sq, occs[i]); } for (int tries = 0; tries < 100000000; tries++) { bitboard_t magic = random_magic(); if (__builtin_popcountll((magic * mask) & 0xFF00000000000000ULL) < 6) continue; memset(used, 0, count * sizeof(bitboard_t)); bool fail = false; for (int i = 0; i < count && !fail; i++) { int idx = (int)(((occs[i] * magic) >> (64 - bits)) & (count - 1)); if (used[idx] == 0) { used[idx] = atts[i]; } else if (used[idx] != atts[i]) { fail = true; } } if (!fail) { entry->magic = magic; for (int i = 0; i < count; i++) { int idx = (int)(((occs[i] * magic) >> (64 - bits)) & (count - 1)); entry->attacks[idx] = atts[i]; } return; } } assert(0 && "Failed to find magic number"); } void magic_init(void) { for (int sq = 0; sq < 64; sq++) { rook_table[sq].mask = compute_rook_mask(sq); rook_table[sq].shift = 64 - __builtin_popcountll(rook_table[sq].mask); bishop_table[sq].mask = compute_bishop_mask(sq); bishop_table[sq].shift = 64 - __builtin_popcountll(bishop_table[sq].mask); } int rook_total = 0, bishop_total = 0; for (int sq = 0; sq < 64; sq++) { rook_total += 1 << (64 - rook_table[sq].shift); bishop_total += 1 << (64 - bishop_table[sq].shift); } rook_attacks_flat = (bitboard_t*)malloc(rook_total * sizeof(bitboard_t)); bishop_attacks_flat = (bitboard_t*)malloc(bishop_total * sizeof(bitboard_t)); assert(rook_attacks_flat && bishop_attacks_flat); int rook_off = 0, bishop_off = 0; for (int sq = 0; sq < 64; sq++) { rook_table[sq].attacks = rook_attacks_flat + rook_off; rook_off += 1 << (64 - rook_table[sq].shift); bishop_table[sq].attacks = bishop_attacks_flat + bishop_off; bishop_off += 1 << (64 - bishop_table[sq].shift); } int max_count = 1 << 12; bitboard_t *occs = (bitboard_t*)malloc(max_count * sizeof(bitboard_t)); bitboard_t *atts = (bitboard_t*)malloc(max_count * sizeof(bitboard_t)); bitboard_t *used = (bitboard_t*)malloc(max_count * sizeof(bitboard_t)); assert(occs && atts && used); for (int sq = 0; sq < 64; sq++) { init_magic_for_square(sq, &rook_table[sq], false, occs, atts, used); } for (int sq = 0; sq < 64; sq++) { init_magic_for_square(sq, &bishop_table[sq], true, occs, atts, used); } free(occs); free(atts); free(used); } bitboard_t get_rook_attacks(int sq, bitboard_t occ) { magic_t *m = &rook_table[sq]; return m->attacks[((occ & m->mask) * m->magic) >> m->shift]; } bitboard_t get_bishop_attacks(int sq, bitboard_t occ) { magic_t *m = &bishop_table[sq]; return m->attacks[((occ & m->mask) * m->magic) >> m->shift]; }