import numpy as np
import matplotlib.pyplot as plt

# Definicija sistema nelinearnih jednačina
def system_equations(x, y):
    eq1 = x**2 + y**2 - 25
    eq2 = x*y - 9
    return eq1, eq2

# Genetski algoritam
def genetic_algorithm(population_size, generations):
    # Inicijalizacija populacije
    population = np.random.uniform(low=-5, high=5, size=(population_size, 2))

    for generation in range(generations):
        # Evaluacija prilagođenosti
        fitness = np.sum(np.abs(np.array(system_equations(population[:, 0], population[:, 1]))), axis=0)

        # Selekcija roditelja
        selected_indices = np.argsort(fitness)[:population_size//2]
        parents = population[selected_indices]

        # Ukrštanje (jednostavno u sredini)
        crossover_point = parents.shape[0] // 2
        children = np.vstack((parents[:crossover_point], parents[crossover_point:][::-1]))

        # Mutacija (dodavanje male slučajne vrednosti)
        mutation_rate = 0.1
        mutation = np.random.uniform(low=-0.1, high=0.1, size=children.shape)
        children += mutation

        # Zamena starih jedinki novim
        population[selected_indices] = children

    # Prikaz rezultata
    best_solution = population[np.argmin(fitness)]
    return best_solution

# Pokretanje genetskog algoritma
best_solution = genetic_algorithm(population_size=50, generations=100)

# Prikaz rezultata
print("The best solution:", best_solution)
print("System values:", system_equations(*best_solution))

# Grafički prikaz sistema
x_vals = np.linspace(-5, 5, 100)
y_vals = np.linspace(-5, 5, 100)
X, Y = np.meshgrid(x_vals, y_vals)
Z1, Z2 = system_equations(X, Y)

plt.contour(X, Y, Z1, levels=[0], colors='r', label='Equation 1')
plt.contour(X, Y, Z2, levels=[0], colors='b', label='Equation 2')
plt.scatter(*best_solution, color='green', marker='*', label='The best solution')
plt.xlabel('x')
plt.ylabel('y')
plt.legend()
plt.title('Solving a system of nonlinear equations with a genetic algorithm')
plt.show()
