Skip to content

N-Queens Problem

import numpy
from deap_er import Fitness, Toolbox, creator, tools

tools.rng.seed(1234)  # disables randomization

BOARD_SIZE = 20


def evaluate(individual):
    size = len(individual)
    left_diagonal = [0] * (2 * size - 1)
    right_diagonal = [0] * (2 * size - 1)

    for i in range(size):
        l_idx = i + individual[i]
        left_diagonal[l_idx] += 1
        r_idx = size - 1 - i + individual[i]
        right_diagonal[r_idx] += 1

    sum_ = 0
    for i in range(2 * size - 1):
        if left_diagonal[i] > 1:
            sum_ += left_diagonal[i] - 1
        if right_diagonal[i] > 1:
            sum_ += right_diagonal[i] - 1

    return (sum_,)  # The comma is essential here.


def setup():
    creator.create_type("FitnessMin", Fitness, weights=(-1.0,))
    creator.create_type("Individual", list, fitness=creator.FitnessMin)

    toolbox = Toolbox()
    toolbox.register("permutation", tools.rng.sample, range(BOARD_SIZE), BOARD_SIZE)
    toolbox.register("individual", tools.init_iterate, creator.Individual, toolbox.permutation)
    toolbox.register("population", tools.init_repeat, list, toolbox.individual)
    toolbox.register("mate", tools.cx_partially_matched)
    toolbox.register("mutate", tools.mut_shuffle_indexes, mut_prob=2.0 / BOARD_SIZE)
    toolbox.register("select", tools.sel_tournament, contestants=3)
    toolbox.register("evaluate", evaluate)

    stats = tools.Statistics(lambda ind: ind.fitness.values)
    stats.register("Avg", numpy.mean)
    stats.register("Std", numpy.std)
    stats.register("Min", numpy.min)
    stats.register("Max", numpy.max)

    return toolbox, stats


def print_results(best_ind):
    if best_ind.fitness.values != (0.0,):
        raise RuntimeError("Evolution failed to converge.")
    print(f"\nRow numbers for each queen on each column of the chessboard: \n{best_ind}")
    print("\nEvolution converged correctly.")


def main():
    toolbox, stats = setup()
    pop = toolbox.population(size=300)
    hof = tools.HallOfFame(1)
    args = {
        "toolbox": toolbox,
        "population": pop,
        "generations": 400,
        "cx_prob": 0.6,
        "mut_prob": 0.3,
        "hof": hof,
        "stats": stats,
        "verbose": True,  # prints stats
    }
    tools.ea_simple(**args)
    print_results(hof[0])


if __name__ == "__main__":
    main()