#exec(open("Gudhi_Circ2.py").read())

import gudhi
#import gudhi.representations
import numpy as np
import matplotlib.pyplot as plt
from scipy.spatial.distance import pdist, squareform
#import warnings
#warnings.filterwarnings('ignore')

from gudhi.wasserstein.barycenter import lagrangian_barycenter as bary
from gudhi.persistence_graphical_tools import plot_persistence_diagram

#from smeUtils import *
#from smeGudhi import *

exec(open("smeUtils.py").read())
exec(open("smeGudhi.py").read())
exec(open("RipsComplex_Aux.py").read())

n_pontos = 80
angulos = np.linspace(0, 2*np.pi, n_pontos, endpoint=False)
raio = 1.0
ruido = 0.1
maxDim = 2

pts = gera2circs(n_pontos, centros=[[-3,0],[3,0]], raio=raio, ruido=ruido)
pts = np.vstack(pts)

analisador = AnalisadorBuracos(pts, dimensao_maxima=maxDim)
analisador.visualizar_pontos()
smeGhudiPlotRipsComplexFromPts(pts, max_edge_length=0.3, title="2 Circunferências", show=True)

simplex_tree = analisador.construir_rips_complex(raio_maximo=0.5)
analisador.persistence = simplex_tree.persistence()
analisador.calcular_persistencia(simplex_tree)

analisador.visualizar_diag_persistencia()
gudhi.plot_persistence_diagram(analisador.persistence)
plt.show()

# Números de Betti:
numBetti = simplex_tree.betti_numbers()
print(f"Números de Betti: {numBetti}")

# Buracos significativos
buracos = analisador.encontrar_buracos_significativos(limiar_persistencia=0.05)
print(f"Buracos significativos encontrados: {len(buracos)}")

#######################################################

pts1 = gera2circs(n_pontos, centros=[[-3,0],[3,0]], raio=raio, ruido=ruido)
pts1 = np.vstack(pts1)

pts2 = gera2circs(n_pontos, centros=[[-3,0],[3,0]], raio=raio, ruido=ruido)
pts2 = np.vstack(pts2)

diags = [pts, pts1, pts2]

# we initialize our estimation on the first diagram (the red one.)
b, log = bary(diags, init=0, verbose=True)

print("Energy reached by this estimation of the barycenter: E=%.2f." %log['energy'])
print("Convergenced made after %s steps." %log['nb_iter'])

G = log["groupings"]

def proj_on_diag(x):
    return ((x[1] + x[0]) / 2, (x[1] + x[0]) / 2)

fig = plt.figure(figsize=(6,6))
ax = fig.add_subplot(111)
colors = ['r', 'b', 'g']

for diag, c in zip(diags, colors):
    plot_persistence_diagram(diag, axes=ax, colormap=c)

def plot_bary(b, diags, groupings, axes):
    # n_y = len(Y.points)
    for i in range(len(diags)):
        indices = G[i]
        n_i = len(diags[i])

        for (y_j, x_i_j) in indices:
            y = b[y_j]
            if y[0] != y[1]:
                if x_i_j >= 0:  # not mapped with the diag
                    x = diags[i][x_i_j]
                else:  # y_j is matched to the diagonal
                    x = proj_on_diag(y)
                ax.plot([y[0], x[0]], [y[1], x[1]], c='black',
                        linestyle="dashed")

    ax.scatter(b[:,0], b[:,1], color='purple', marker='d', label="barycenter (estim)")
    ax.legend()
    ax.set_title("Set of diagrams and their barycenter", fontsize=22)

plot_bary(b, diags, G, axes=ax)

fig, axs = plt.subplots(1, 3, figsize=(15, 5))

colors = ['r', 'b', 'g']

for i, ax in enumerate(axs):
    for diag, c in zip(diags, colors):
        plot_persistence_diagram(diag, axes=ax, colormap=c)

    b, log = bary(diags, init=i, verbose=True)
    e = log["energy"]
    G = log["groupings"]
    # print(G)
    plot_bary(b, diags, groupings=G, axes=ax)
    ax.set_title("Barycenter estim with init=%s. Energy: %.2f" %(i, e), fontsize=14)
    
plt.show()
