import torch
from torchvision import transforms
import matplotlib.pyplot as plt
from torchvision.transforms import v2
import numpy as np
import cv2
import random

# Definir una composición de transformaciones
transform = transforms.Compose([v2.RandomRotation(degrees=(0, 360)),
                         v2.CenterCrop(size=58),
                         v2.RandomHorizontalFlip(p=0.5),
                         v2.RandomVerticalFlip(p=0.5),
                         #Resize /2
                        ])

def volteo_horizontal(imagen):
    return cv2.flip(imagen, 1) # 1 para volteo horizontal



def volteo_vertical(imagen):
    return cv2.flip(imagen, 0) # 0 para volteo vertical

def escalado(imagen, factor_escala):
    alto, ancho = imagen.shape[:2]
    nuevo_ancho, nuevo_alto = int(ancho * factor_escala), int(alto * factor_escala)
    imagen_escalada = cv2.resize(imagen, (nuevo_ancho, nuevo_alto), interpolation=cv2.INTER_LINEAR)

    # Si se hace zoom in, es posible que quieras recortar al tamaño original
    if factor_escala > 1:
        x = (nuevo_ancho - ancho) // 2
        y = (nuevo_alto - alto) // 2
        imagen_escalada = imagen_escalada[y:y+alto, x:x+ancho]
    return imagen_escalada

def rotacion(imagen, angulo):
    alto, ancho = imagen.shape[:2]
    punto_central = (ancho // 2, alto // 2)
    matriz_rotacion = cv2.getRotationMatrix2D(punto_central, angulo, 1.0)
    imagen_rotada = cv2.warpAffine(imagen, matriz_rotacion, (ancho, alto))
    return imagen_rotada

def cambiar_brillo(imagen, factor_brillo):
    hsv = cv2.cvtColor(imagen, cv2.COLOR_BGR2HSV)
    h, s, v = cv2.split(hsv)
    v = np.clip(v * factor_brillo, 0, 255).astype(np.uint8)
    hsv_brillo = cv2.merge((h, s, v))
    return cv2.cvtColor(hsv_brillo, cv2.COLOR_HSV2BGR)

def transformacion_aleatoria(imagen):
    imagen_transformada = imagen.copy()
    if random.random() > 0.5:
        imagen_transformada = volteo_horizontal(imagen_transformada)
    if random.random() > 0.5:
        imagen_transformada = volteo_vertical(imagen_transformada)

    if random.random() > 0.5:
        imagen_transformada = rotacion(imagen_transformada, random.uniform(-15, 15))
        imagen_transformada = escalado( imagen_transformada, random.uniform(1.12,1.18) )
        '''
    if random.random() > 0.5:
        imagen_transformada = escalado( imagen_transformada, random.uniform(1.1,1.2) )
        '''
    ''' 
    if random.random() > 0.5:
        factor_brillo = random.uniform(0.7, 1.3)
        imagen_transformada = cambiar_brillo(imagen_transformada, factor_brillo)
    '''   
    return imagen_transformada                    

# Cargar una imagen de ejemplo (reemplaza con tu propia imagen)
# leermos las etiquetas
desEtiquetas = np.loadtxt("etiquetas.csv",
                 delimiter=",", dtype=str)

def busca_etiqueta_real(imagen):
    for i in range(len(desEtiquetas)):
        if desEtiquetas[i][0] == str(str(imagen)+".png"):
            return int(desEtiquetas[i][1]) # retorna la etiqueta de la tupla leida
aumento = 200

contador = 0
etiquetasOriginales = []
nTitulos = []
for i in range(480):
    for a in range(aumento):
        ruta = "T/"+str(i)+".png"
        etiquetaOrigin = busca_etiqueta_real(i)
        etiquetasOriginales.append(etiquetaOrigin)
        nu = cv2.imread( ruta )
        nu = transformacion_aleatoria(nu)
        #nu = ruido.agregar_ruido_gaussianoST1C(nu, medias[m], desEst[d])
		titulo = "aa/" + str(contador)
		cv2.imwrite(f"{titulo}.png", nu)
        nTitulos.append(str(contador)+".png")
        classes = ['Cruz', 'Estrella', 'I', 'Linea', 'Ovalo', 'Pentagono', 'Rectangulo', 'Triangulo']
        contador += 1

for i in range(contador):
    nuevoTexto = nTitulos[i]+", "+str(etiquetasOriginales[i])+"\n"
    f = open("etiquetas_aumentado.csv", "a")
    f.write(nuevoTexto)
    f.close()


'''


# Aplicar la transformación a la imagen
transformed_img = transform(img)

# Visualizar la imagen original y la transformada (opcional)
fig, axes = plt.subplots(1, 2)
axes[0].imshow(img)
axes[0].set_title("Original")
axes[0].axis('off')

# Para visualizar el tensor, necesitamos cambiar las dimensiones y desnormalizar si se aplicó alguna normalización
def show_tensor_image(tensor):
    permuted_tensor = tensor.permute(1, 2, 0) # Cambiar de (C, H, W) a (H, W, C)
    axes[1].imshow(permuted_tensor)
    axes[1].set_title("Transformada")
    axes[1].axis('off')

show_tensor_image(transformed_img)
plt.show()

'''
