#! /usr/bin/env python
# -*- coding: utf-8 -*-

"""
my_de.py
Implementing the differential evolution algorithm.

Ernesto Costa - 25 de Março 2010.
"""

# ---------------  TODO: still lots of things!!!

import matplotlib
from pylab import *
from random import random,randint, shuffle, uniform,sample, choice
from operator import itemgetter
from math import sqrt, sin, cos, pi
from copy import copy, deepcopy

def run(num_runs,num_gera,tam_pop, domain, prob_cruza, gama, fit_func):
    # Colecta Dados
    print 'Wait, please '
    estatistica_total = [ga(num_gera,tam_pop,domain, prob_cruza,gama, fit_func) for i in range(num_runs)]
    print "That's it!"
    # Processa Dados: melhor e médias por geração
    resultados_gera = zip(*estatistica_total)   
    melhores = [max([indiv[0] for indiv in gera])for gera in resultados_gera]
    medias = [sum([indiv[1] for indiv in gera])/float(num_runs) for gera in resultados_gera]
    # Mostra
    ylabel('Fitness')
    xlabel('Generation')
    titulo = 'Differential Evolution --  Gamma: %0.2f , Crossover: %0.2f' % (gama,prob_cruza)
    title(titulo)
    axis= [0,num_gera,0,len(domain)]
    p1 = plot(melhores,'r-o',label="Best")
    p2 = plot(medias,'g-s',label="Average")
    legend(loc='lower right')
    show()



def ga(num_gera,tam_pop, domain, prob_cruza,gama, fit_func):
    # Estatistica : lista de pares (melhor, média)
    # Podia ter usado um ficheiro...
    tam_cromo = len(domain)
    estatistica=[]
    # inicializa população: indiv = (cromo,fit)
    populacao = [[gera_indiv(domain),0] for j in range(tam_pop)]

    

    # avalia população
    populacao = [ [indiv[0], fit_func(indiv[0])] for indiv in populacao]
    populacao.sort(key=itemgetter(1),reverse=True) 

    while num_gera:
        # Estatistica
        qualidade_media = sum([indiv[1] for indiv in populacao])/ float(tam_pop)
        qualidade_melhor = populacao[0][1]
        estatistica.append((qualidade_melhor,qualidade_media))
	# Descendentes
	filhos = []
	for  elem in populacao:

	    mutante = deepcopy(elem)
	    # mutação 
	    # selecciona 3 indivíduos
	    pool = sample(populacao,3)

	    while mutante in pool:
		pool = sample(populacao,3)

	    # selecciona gene
	    index = choice(range(tam_cromo))
	    # altera
	    for i in range(tam_cromo):
		if (i == index) or (random() < prob_cruza):
		    mutante[0][index] = pool[0][0][i] + gama * (pool[1][0][i] - pool[2][0][i])
	    filhos.append(mutante)
	    
	# avalia mutantes	    
	filhos = [ [indiv[0], fit_func(indiv[0])] for indiv in filhos]
        # nova população
	populacao = sobreviventes(populacao,filhos)
	
        # Avalia e ordena nova _população
        populacao = [[indiv[0], fit_func(indiv[0])] for indiv in populacao]
        populacao.sort(key=itemgetter(1),reverse=True) 
        num_gera = num_gera - 1
    print "Melhor: \n%s\n%s" % (populacao[0][0], populacao[0][1])
    return estatistica


# Representação 


def gera_indiv(domain):
    indiv = [uniform(inf,sup) for inf,sup in domain]
    return indiv


# sobreviventes
def sobreviventes(pais,filhos):
    """ Admite que não há problemas de tamanhos... Problema de maximização"""
    tam = len(pais)
    novos_pais = []
    for i in range(tam):
	if pais[i][1] > filhos[i][1]:
	    novos_pais.append(pais[i])
	else:
	    novos_pais.append(filhos[i])	    
    return novos_pais



# Qualidade: 

# De Jong

def de_jong_f1(indiv):
    """ domínio [[-5.12,5.12],[-5.12,5.12],[-5.12,5.12]]"""
    # validate values. Clip if outside bounds.
    x = indiv[0]
    y = indiv[1]
    z = indiv[2]
    if (x < -5.12) or (x > 5.12) or (y < -5.12) or (y > 5.12) or (z < -5.12) or (z > 5.12):
	    return 0
    else:
	f = indiv[0]**2 + indiv[1]**2 + indiv[2]**2
	return (1.0/(1 + f))

    
def de_jong_f2(indiv):
    """ domínio [[-2.048,2.048],[-2.048,2.048]]"""
    # validate values. Clip if outside bounds.
    x = indiv[0]
    y = indiv[1]
    if (x < -2.048) or (x > 2.048) or (y < -2.048) or (y >2.048):
	    return 0
    else:
	f = 100* (indiv[0]**2 - indiv[1])**2 + (1 - indiv[0])**2
	return (1.0/(1 + f))

# -- Michaelewicz	
def michalewicz_1(indiv):
	x=indiv[0]
	if (x < -1) or (x >2):
	    return 0
	else:
	    res= x * sin(10 * pi * x) + 1.0
	    return res
	

def michalewicz_2(indiv):
    """ Máximo = 38.85."""
    x=indiv[0]
    y=indiv[1]
    if (x < -3.0) or (x >12.1) or (y < 4.1) or (y > 5.8):
	return 0
    else:
	res= x * sin(4 * pi * x) + y* sin(20 *pi * y) + 21.5
	return res
    
def rastringin_3d(indiv):
	""" Domínio [(-5.12,5.12),(-5.12,5.12),(-5.12,5.12)]."""
	# validate values. Clip if outside bounds.
	x=indiv[0]
	y = indiv[1]
	z = indiv[2]
	if (x < -5.12) or (x > 5.12) or (y < -5.12) or (y > 5.12) or (z < -5.12) or (z > 5.12):
	    return 0
	else:
	    f = 3 * 10.0 + (x**2 - 10.0 * cos(2*pi*x)) + (y**2 - 10.0 * cos(2*pi*y)) + (z**2 - 10.0 * cos(2*pi*z))
	    return f  
    

if __name__ == '__main__':
    """
    domain = [(-1,2)]
    print gera_indiv(domain)
    """
    #run(3,20,30,[[-5.12,5.12],[-5.12,5.12],[-5.12,5.12]],0.8,0.5,de_jong_f1)
    #run(10,60,30,[[-2.048,2.048],[-2.048,2.048]],0.8,0.5,de_jong_f2)
    #run(10,60,30,[[-3.0,12.2],[4.1,5.8]],0.8,0.5,michalewicz_2)
    run(10,60,30,[[-5.12,5.12],[-5.12,5.12],[-5.12,5.12]],0.8,0.5,rastringin_3d)
    
