#pyright: basic
import pprint
import random

POP_SIZE = 50
MAX_GEN = 100
DIM = 25
MUT_PROB = 0.2
MUT_FLIP_PROB = 1/DIM
CROSS_PROB = 0.8

def fitness(x):
    return sum(x)

def create_random_individual():
    return [random.randint(0,1) for _ in range(DIM)]

def create_random_population():
    return [create_random_individual() for _ in range(POP_SIZE)]

def select(pop, fits):
    return random.choices(pop, weights=fits, k=POP_SIZE)

def cross(p1, p2):
    point = random.randrange(DIM)
    o1 = p1[:point] + p2[point:]
    o2 = p2[:point] + p1[point:]
    return o1, o2

def crossover(pool):
    o = []
    for (p1, p2) in zip(pool[::2], pool[1::2]):
        o1, o2 = p1[:], p2[:]
        if random.random() < CROSS_PROB:
            o1, o2 = cross(p1, p2)
        o.extend([o1, o2])
    return o

def mutate(ind):
    return [1 - i if random.random() < MUT_FLIP_PROB else i for i in ind]
    # out = []
    # for i in ind:
    #    if random.random() < MUT_FLIP_PROB:
    #        out.append(1-i)
    #    else:
    #        out.append(i)
    # return o

def mutation(pool):
    return [mutate(ind) if random.random() < MUT_PROB else ind[:] for ind in pool]

def evolutionary_algorithm(elitism=False):
    pop = create_random_population()
    log = []
    for G in range(MAX_GEN):
        fits = [fitness(x) for x in pop]
        log.append(max(fits)) 
        mating_pool = select(pop, fits)
        o = crossover(mating_pool)
        off = mutation(o)
        if elitism:
            pop = off[1:] + [max(pop, key=fitness)]
        else:
            pop = off
    return pop, log

logs_noel = []
for _ in range(100):
    pop, log = evolutionary_algorithm(elitism=False)
    logs_noel.append(log)

logs_el = []
for _ in range(100):
    pop, log = evolutionary_algorithm(elitism=False)
    logs_el.append(log)

import matplotlib.pyplot as plt 
import numpy as np


x = list(range(MAX_GEN))
logs_noel = np.array(logs_noel)
plt.plot(logs_noel.mean(axis=0))

low = np.percentile(logs_noel, q=25, axis=0)
high = np.percentile(logs_noel, q=75, axis=0)

plt.fill_between(x, low, high, alpha=0.5)

logs_el = np.array(logs_el)
plt.plot(logs_el.mean(axis=0))
low = np.percentile(logs_el, q=25, axis=0)
high = np.percentile(logs_el, q=75, axis=0)

plt.fill_between(x, low, high, alpha=0.5)
plt.show()
