import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets
from torchvision.transforms import ToTensor, Resize
import torch.nn.functional as F
from torch.utils.data import Dataset
import pandas as pd
import os
import cv2
from torchvision import transforms
from torchvision.transforms import v2
import matplotlib.pyplot as plt
import math

from functools import partial
from typing import Any, Callable, List, Optional
import warnings
from typing import Any, Dict, List, Optional
from torch import Tensor
from torchvision.transforms._presets import ImageClassification
from torchvision.utils import _log_api_usage_once
from torchvision.models._api import register_model, Weights, WeightsEnum
from torchvision.models._meta import _IMAGENET_CATEGORIES
from torchvision.models._utils import _ovewrite_named_param, handle_legacy_interface

from collections import namedtuple
from typing import Any, Callable, List, Optional, Tuple
import numpy as np

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

def AumentaDatos(archivos, num):
    for archivo in archivos:
        archivo = path+archivo
        df = pd.read_csv(archivo, header=None)
        aumentados = df.values.tolist()
        aumentados *= num
        # print('aumentados_' + archivo[archivo.index('_')+1:])
        np.savetxt('aumentados_' + archivo[archivo.index('_')+1:], aumentados, delimiter =', ', fmt='%s')

class CustomImageDataset(Dataset):
    def __init__(self, annotations_file, img_dir, transform=None, target_transform=None):
        self.img_labels = pd.read_csv(annotations_file)
        self.img_dir = img_dir
        self.transform = transform
        self.target_transform = target_transform

    def __len__(self):
        return len(self.img_labels)

    def __getitem__(self, idx):
        img_path = os.path.join(self.img_dir, self.img_labels.iloc[idx, 0])
        #image = read_image(img_path)
        image = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)  # verficar maximo y minimo Opencv normalizador
        image = cv2.normalize(image, image, 0, 1.0, cv2.NORM_MINMAX) #Normaliza con OpenCV
        label = self.img_labels.iloc[idx, 1]
        if self.transform:
            image = self.transform(image)
        if self.target_transform:
            label = self.target_transform(label)
        return image, label


#Transformaciones para las imagenes
# transforma = transforms.Compose([transforms.ToTensor()])
transforma = v2.Compose([v2.RandomRotation(degrees=(0, 360)),
                         v2.CenterCrop(size=58),
                         v2.RandomHorizontalFlip(p=0.5),
                         v2.RandomVerticalFlip(p=0.5),
                         #Resize /2
                         transforms.ToTensor()]) #ToTensor


path = '/home/b/Documentos/TESIS/R49 EfficientNet GoogLeNet x200/'
archivos = ["labels_Train.csv", "labels_Test.csv"]
aumenta = 200
AumentaDatos(archivos, aumenta) #Cantidad de veces que se aumentaran los datos

#Crea los Sets de Datos
train_data = CustomImageDataset(path+'aumentados_Train.csv', path+'BDF_pytorch/Train/', transforma)
test_data = CustomImageDataset(path+'aumentados_Test.csv', path+'BDF_pytorch/Test/', transforma)


#Tamaño del batch
batch_size = 32


#Crea los Data Loaders
train_dataloader = DataLoader(train_data, batch_size=batch_size, shuffle=True, drop_last=True, num_workers=4)
test_dataloader = DataLoader(test_data, batch_size=batch_size, shuffle=True, drop_last=True, num_workers=4)


# Para calcular el tamaño del salida
#      (ne - k + 2p)
# ns = ------------- + 1
#            s
# Donde: ns es el tamaño de salida
#        ne es el tamaño de entrada
#        k es el tamaño del kernel
#        p es el padding que se le aplica a la imagen
#        s es el stride


# Define model
__all__ = ["GoogLeNet", "GoogLeNetOutputs", "_GoogLeNetOutputs", "GoogLeNet_Weights", "googlenet"]


GoogLeNetOutputs = namedtuple("GoogLeNetOutputs", ["logits", "aux_logits2", "aux_logits1"])
GoogLeNetOutputs.__annotations__ = {"logits": Tensor, "aux_logits2": Optional[Tensor], "aux_logits1": Optional[Tensor]}

# Script annotations failed with _GoogleNetOutputs = namedtuple ...
# _GoogLeNetOutputs set here for backwards compat
_GoogLeNetOutputs = GoogLeNetOutputs


class GoogLeNet(nn.Module):
    __constants__ = ["aux_logits", "transform_input"]

    def __init__(
        self,
        num_classes: int = 8,
        aux_logits: bool = True,
        transform_input: bool = False,
        init_weights: Optional[bool] = None,
        blocks: Optional[List[Callable[..., nn.Module]]] = None,
        dropout: float = 0.2,
        dropout_aux: float = 0.7,
    ) -> None:
        super().__init__()
        _log_api_usage_once(self)
        if blocks is None:
            blocks = [BasicConv2d, Inception, InceptionAux]
        if init_weights is None:
            warnings.warn(
                "The default weight initialization of GoogleNet will be changed in future releases of "
                "torchvision. If you wish to keep the old behavior (which leads to long initialization times"
                " due to scipy/scipy#11299), please set init_weights=True.",
                FutureWarning,
            )
            init_weights = True
        if len(blocks) != 3:
            raise ValueError(f"blocks length should be 3 instead of {len(blocks)}")
        conv_block = blocks[0]
        inception_block = blocks[1]
        inception_aux_block = blocks[2]

        self.aux_logits = aux_logits
        self.transform_input = transform_input

        self.conv1 = conv_block(1, 64, kernel_size=7, stride=2, padding=3)
        self.maxpool1 = nn.MaxPool2d(3, stride=2, ceil_mode=True)
        self.conv2 = conv_block(64, 64, kernel_size=1)
        self.conv3 = conv_block(64, 192, kernel_size=3, padding=1)
        self.maxpool2 = nn.MaxPool2d(3, stride=2, ceil_mode=True)

        self.inception3a = inception_block(192, 64, 96, 128, 16, 32, 32)
        # self.inception3b = inception_block(256, 128, 128, 192, 32, 96, 64)
        # self.maxpool3 = nn.MaxPool2d(3, stride=2, ceil_mode=True)

        # self.inception4a = inception_block(480, 192, 96, 208, 16, 48, 64)
        # self.inception4b = inception_block(512, 160, 112, 224, 24, 64, 64)
        # self.inception4c = inception_block(512, 128, 128, 256, 24, 64, 64)
        # self.inception4d = inception_block(512, 112, 144, 288, 32, 64, 64)
        # self.inception4e = inception_block(528, 256, 160, 320, 32, 128, 128)
        # self.maxpool4 = nn.MaxPool2d(2, stride=2, ceil_mode=True)

        # self.inception5a = inception_block(832, 256, 160, 320, 32, 128, 128)
        # self.inception5b = inception_block(832, 384, 192, 384, 48, 128, 128)

        if aux_logits:                     #512
            self.aux1 = inception_aux_block(256, num_classes, dropout=dropout_aux)
                                           #528
            self.aux2 = inception_aux_block(256, num_classes, dropout=dropout_aux)
        else:
            self.aux1 = None  # type: ignore[assignment]
            self.aux2 = None  # type: ignore[assignment]

        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
        self.dropout = nn.Dropout(p=dropout)
        self.fc = nn.Linear(256, num_classes)

        if init_weights:
            for m in self.modules():
                if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear):
                    torch.nn.init.trunc_normal_(m.weight, mean=0.0, std=0.01, a=-2, b=2)
                elif isinstance(m, nn.BatchNorm2d):
                    nn.init.constant_(m.weight, 1)
                    nn.init.constant_(m.bias, 0)

    def _transform_input(self, x: Tensor) -> Tensor:
        if self.transform_input:
            x_ch0 = torch.unsqueeze(x[:, 0], 1) * (0.229 / 0.5) + (0.485 - 0.5) / 0.5
            x_ch1 = torch.unsqueeze(x[:, 1], 1) * (0.224 / 0.5) + (0.456 - 0.5) / 0.5
            x_ch2 = torch.unsqueeze(x[:, 2], 1) * (0.225 / 0.5) + (0.406 - 0.5) / 0.5
            x = torch.cat((x_ch0, x_ch1, x_ch2), 1)
        return x

    def _forward(self, x: Tensor) -> Tuple[Tensor, Optional[Tensor], Optional[Tensor]]:
        # N x 3 x 224 x 224
        x = self.conv1(x)
        # N x 64 x 112 x 112
        x = self.maxpool1(x)
        # N x 64 x 56 x 56
        x = self.conv2(x)
        # N x 64 x 56 x 56
        x = self.conv3(x)
        # N x 192 x 56 x 56
        x = self.maxpool2(x)

        # N x 192 x 28 x 28
        x = self.inception3a(x)
        # N x 256 x 28 x 28
        # x = self.inception3b(x)
        # N x 480 x 28 x 28
        # x = self.maxpool3(x)
        # N x 480 x 14 x 14
        # x = self.inception4a(x)
        # N x 512 x 14 x 14
        aux1: Optional[Tensor] = None
        if self.aux1 is not None:
            if self.training:
                aux1 = self.aux1(x)

        # x = self.inception4b(x)
        # N x 512 x 14 x 14
        # x = self.inception4c(x)
        # N x 512 x 14 x 14
        # x = self.inception4d(x)
        # N x 528 x 14 x 14
        aux2: Optional[Tensor] = None
        if self.aux2 is not None:
            if self.training:
                aux2 = self.aux2(x)

        # x = self.inception4e(x)
        # N x 832 x 14 x 14
        # x = self.maxpool4(x)
        # N x 832 x 7 x 7
        # x = self.inception5a(x)
        # N x 832 x 7 x 7
        # x = self.inception5b(x)
        # N x 1024 x 7 x 7

        x = self.avgpool(x)
        # N x 1024 x 1 x 1
        x = torch.flatten(x, 1)
        # N x 1024
        x = self.dropout(x)
        x = self.fc(x)
        # N x 1000 (num_classes)
        return x, aux2, aux1

    # @torch.jit.unused
    def eager_outputs(self, x: Tensor, aux2: Tensor, aux1: Optional[Tensor]) -> GoogLeNetOutputs:
        if self.training and self.aux_logits:
            return _GoogLeNetOutputs(x, aux2, aux1)
        else:
            return x  # type: ignore[return-value]

    def forward(self, x: Tensor) -> GoogLeNetOutputs:
        x = self._transform_input(x)
        x, aux1, aux2 = self._forward(x)
        aux_defined = self.training and self.aux_logits
        if torch.jit.is_scripting():
            if not aux_defined:
                warnings.warn("Scripted GoogleNet always returns GoogleNetOutputs Tuple")
            return GoogLeNetOutputs(x, aux2, aux1)
        else:
            return self.eager_outputs(x, aux2, aux1)


class Inception(nn.Module):
    def __init__(
        self,
        in_channels: int,
        ch1x1: int,
        ch3x3red: int,
        ch3x3: int,
        ch5x5red: int,
        ch5x5: int,
        pool_proj: int,
        conv_block: Optional[Callable[..., nn.Module]] = None,
    ) -> None:
        super().__init__()
        if conv_block is None:
            conv_block = BasicConv2d
        self.branch1 = conv_block(in_channels, ch1x1, kernel_size=1)

        self.branch2 = nn.Sequential(
            conv_block(in_channels, ch3x3red, kernel_size=1), conv_block(ch3x3red, ch3x3, kernel_size=3, padding=1)
        )

        self.branch3 = nn.Sequential(
            conv_block(in_channels, ch5x5red, kernel_size=1),
            # Here, kernel_size=3 instead of kernel_size=5 is a known bug.
            # Please see https://github.com/pytorch/vision/issues/906 for details.
            conv_block(ch5x5red, ch5x5, kernel_size=3, padding=1),
        )

        self.branch4 = nn.Sequential(
            nn.MaxPool2d(kernel_size=3, stride=1, padding=1, ceil_mode=True),
            conv_block(in_channels, pool_proj, kernel_size=1),
        )

    def _forward(self, x: Tensor) -> List[Tensor]:
        branch1 = self.branch1(x)
        branch2 = self.branch2(x)
        branch3 = self.branch3(x)
        branch4 = self.branch4(x)

        outputs = [branch1, branch2, branch3, branch4]
        return outputs

    def forward(self, x: Tensor) -> Tensor:
        outputs = self._forward(x)
        return torch.cat(outputs, 1)


class InceptionAux(nn.Module):
    def __init__(
        self,
        in_channels: int,
        num_classes: int,
        conv_block: Optional[Callable[..., nn.Module]] = None,
        dropout: float = 0.7,
    ) -> None:
        super().__init__()
        if conv_block is None:
            conv_block = BasicConv2d
        self.conv = conv_block(in_channels, 128, kernel_size=1)

        self.fc1 = nn.Linear(2048, 1024)
        self.fc2 = nn.Linear(1024, num_classes)
        self.dropout = nn.Dropout(p=dropout)

    def forward(self, x: Tensor) -> Tensor:
        # aux1: N x 512 x 14 x 14, aux2: N x 528 x 14 x 14
        x = F.adaptive_avg_pool2d(x, (4, 4))
        # aux1: N x 512 x 4 x 4, aux2: N x 528 x 4 x 4
        x = self.conv(x)
        # N x 128 x 4 x 4
        x = torch.flatten(x, 1)
        # N x 2048
        x = F.relu(self.fc1(x), inplace=True)
        # N x 1024
        x = self.dropout(x)
        # N x 1024
        x = self.fc2(x)
        # N x 1000 (num_classes)

        return x


class BasicConv2d(nn.Module):
    def __init__(self, in_channels: int, out_channels: int, **kwargs: Any) -> None:
        super().__init__()
        self.conv = nn.Conv2d(in_channels, out_channels, bias=False, **kwargs)
        self.bn = nn.BatchNorm2d(out_channels, eps=0.001)

    def forward(self, x: Tensor) -> Tensor:
        x = self.conv(x)
        x = self.bn(x)
        return F.relu(x, inplace=True)


class GoogLeNet_Weights(WeightsEnum):
    IMAGENET1K_V1 = Weights(
        url="https://download.pytorch.org/models/googlenet-1378be20.pth",
        transforms=partial(ImageClassification, crop_size=224),
        meta={
            "num_params": 6624904,
            "min_size": (15, 15),
            "categories": _IMAGENET_CATEGORIES,
            "recipe": "https://github.com/pytorch/vision/tree/main/references/classification#googlenet",
            "_metrics": {
                "ImageNet-1K": {
                    "acc@1": 69.778,
                    "acc@5": 89.530,
                }
            },
            "_ops": 1.498,
            "_file_size": 49.731,
            "_docs": """These weights are ported from the original paper.""",
        },
    )
    DEFAULT = IMAGENET1K_V1


# @register_model()
# @handle_legacy_interface(weights=("pretrained", GoogLeNet_Weights.IMAGENET1K_V1))
def googlenet(*, weights: Optional[GoogLeNet_Weights] = None, progress: bool = True, **kwargs: Any) -> GoogLeNet:
    """GoogLeNet (Inception v1) model architecture from
    `Going Deeper with Convolutions <http://arxiv.org/abs/1409.4842>`_.

    Args:
        weights (:class:`~torchvision.models.GoogLeNet_Weights`, optional): The
            pretrained weights for the model. See
            :class:`~torchvision.models.GoogLeNet_Weights` below for
            more details, and possible values. By default, no pre-trained
            weights are used.
        progress (bool, optional): If True, displays a progress bar of the
            download to stderr. Default is True.
        **kwargs: parameters passed to the ``torchvision.models.GoogLeNet``
            base class. Please refer to the `source code
            <https://github.com/pytorch/vision/blob/main/torchvision/models/googlenet.py>`_
            for more details about this class.
    .. autoclass:: torchvision.models.GoogLeNet_Weights
        :members:
    """
    weights = GoogLeNet_Weights.verify(weights)

    original_aux_logits = kwargs.get("aux_logits", False)
    if weights is not None:
        if "transform_input" not in kwargs:
            _ovewrite_named_param(kwargs, "transform_input", True)
        _ovewrite_named_param(kwargs, "aux_logits", False) #True
        _ovewrite_named_param(kwargs, "init_weights", False)
        _ovewrite_named_param(kwargs, "num_classes", len(weights.meta["categories"]))

    model = GoogLeNet(**kwargs)

    if weights is not None:
        model.load_state_dict(weights.get_state_dict(progress=progress, check_hash=True))
        if not original_aux_logits:
            model.aux_logits = False
            model.aux1 = None  # type: ignore[assignment]
            model.aux2 = None  # type: ignore[assignment]
        else:
            warnings.warn(
                "auxiliary heads in the pretrained googlenet model are NOT pretrained, so make sure to train them"
            )

    return model


def F_checkpoint(model, filename):
    torch.save(model.state_dict(), filename)


nuevo_modelo = googlenet()
print(nuevo_modelo)

# # Imprimir el state_dict del modelo
# print("Model's state_dict:")
# for param_tensor in nuevo_modelo.state_dict():
#     print(param_tensor, "\t", nuevo_modelo.state_dict()[param_tensor].size())


loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(nuevo_modelo.parameters(), lr=0.1, momentum=0.9)
# optimizer = torch.optim.Adadelta( nuevo_modelo.parameters(), lr=0.001, rho=0.9, eps=1e-06, weight_decay=0)
# optimizer = torch.optim.Adam(nuevo_modelo.parameters(), lr=0.001)


# # Se lee el modelo ya entrenado
# checkpoint = torch.load("nm_GoogLeNet.pth")
# nuevo_modelo.load_state_dict( checkpoint )

nuevo_modelo.to(device)
loss_fn.to(device)

# i = 1
# for param in nuevo_modelo.linear_relu_stack.parameters():
#  	param.requires_grad = False
#  	# print( i, param.requires_grad )
#  	# print(type(param), param.size())
#  	i += 1

i = 1
for param in nuevo_modelo.parameters():
 	param.requires_grad = True
 	print( i, param.requires_grad )
 	print(type(param), param.size())
 	i += 1

'''
def train(dataloader, model, loss_fn, optimizer):
    global exactitud_entrena
    size = len(dataloader.dataset)
    model.train()
    train_loss, correct = 0, 0
    for batch, (X, y) in enumerate(dataloader):
        X, y = X.to(device), y.to(device)
        
		# Compute prediction error
        with torch.cuda.amp.autocast(): #Para evitar un desbordamiento numérico en el cálculo de la entropía cruzada
            pred = model(X) #Devuelve un GoogLeOutput
            pred = pred.logits
            loss = loss_fn(pred, y)
        train_loss += loss_fn(pred, y).item()
        correct += (pred.argmax(1) == y).type(torch.float).sum().item()
        
		# Backpropagation
        loss.backward()
        nn.utils.clip_grad_norm_(model.parameters(), 1.0) #Gradient clipp
        optimizer.step()
        optimizer.zero_grad()
        
        if batch % 100 == 0:    
            loss, current = loss.item(), (batch + 1) * len(X)
            print(f"loss: {loss:>7f}  [{current:>5d}/{size:>5d}]\n")
    
    train_loss /= size
    correct /= size
    print(f"Train Error: \n Accuracy: {(100*correct):>0.1f}%, Avg loss: {train_loss:>8f} \n")
    exactitud_entrena.append(correct * 100)
    error_entrena.append(train_loss)
    

def test(dataloader, model, loss_fn):
    global exactitud_prueba, best_accuracy, t, best_epoch
    size = len(dataloader.dataset)
    num_batches = len(dataloader)
    model.eval()
    test_loss, correct = 0, 0
    with torch.no_grad():
        for X, y in dataloader:
            X, y = X.to(device), y.to(device)
            pred = model(X)
            
            test_loss += loss_fn(pred, y).item()
            correct += (pred.argmax(1) == y).type(torch.float).sum().item()
            
    test_loss /= num_batches
    correct /= size
    print(f"Test Error: \n Accuracy: {(100*correct):>0.1f}%, Avg loss: {test_loss:>8f} \n\n")
    exactitud_prueba.append(correct*100)
    error_prueba.append(test_loss)
    
    if (100*correct) > best_accuracy:
        best_accuracy = 100*correct
        best_epoch = t
        F_checkpoint(model, "nm_GoogLeNet.pth")
    # elif t - best_epoch > early_stop:
    #     print("Entrenamiento detenido temprano en la época %d" % t)
    #     return -1  # Termina el ciclo de entrenamiento


# early_stop = 10
best_accuracy = -1
best_epoch = -1


exactitud_entrena = []
exactitud_prueba = []
error_entrena = []
error_prueba = []

epochs = 100
for t in range(epochs):
    print(f"Epoch {t+1}\n-------------------------------")
    if train(train_dataloader, nuevo_modelo, loss_fn, optimizer) == -1: break
    es = test(test_dataloader, nuevo_modelo, loss_fn)
    # if es == -1: break
print("Done!")

print("Mejor exactitud %.2f" % best_accuracy, "% en la epoca", best_epoch)
torch.save(nuevo_modelo.state_dict(), "nm_GoogLeNet_"+str(best_accuracy)[:2]+str(best_accuracy)[3:5]+".pth")
print("Saved PyTorch Model State to " + "nm_GoogLeNet_" + str(best_accuracy)[:2] + str(best_accuracy)[3:5]+".pth")


#Graficas de exactitud y error
muestras = list(range(len(exactitud_entrena)))

fig, ax = plt.subplots(ncols=2, nrows=1)
ax[0].plot(muestras, exactitud_entrena, label='Exactitud entrenamiento')
ax[0].plot(muestras, exactitud_prueba, label='Exactitud prueba')
ax[0].axvline(best_epoch, linestyle='--', color='r',label='Mejor modelo: %.2f' % best_accuracy)
ax[0].scatter(0, exactitud_entrena[0], c='blue', s=50)
ax[0].scatter(0, exactitud_prueba[0], c='orange', s=50)


ax[1].plot(muestras, error_entrena, label='Error entrenamiento')
ax[1].plot(muestras, error_prueba, label='Error prueba')
min_error = error_prueba.index(min(error_prueba))
ax[1].axvline(best_epoch, linestyle='--', color='r',label='Minimo error prueba: epo'+ str(best_epoch))


ax[0].title.set_text("Exactitud")
ax[1].title.set_text("Error")
ax[0].legend()
ax[1].legend()

# ax[0].set_xlim(0, epochs)
ax[0].set_ylim(0, 100)
# ax[1].set_xlim(0, epochs)
ax[1].set_ylim(0, 1.0)
#Guardar la imagen
plt.show()

plt.savefig('GoogLeNet_'+str(best_accuracy)[:2]+str(best_accuracy)[3:5]+'.png')

# liberamos memoria de GPU
import gc
gc.collect()
torch.cuda.empty_cache()
with torch.no_grad():
    torch.cuda.empty_cache()
'''

import torchvision
from torchvision.utils import save_image

#Para ver que imágenes no están bien clasificadas
checkpoint = torch.load("nm_GoogLeNet_9987.pth")
nuevo_modelo.load_state_dict( checkpoint )
nuevo_modelo.eval()

train_dataloader = DataLoader(train_data, batch_size=1, shuffle=True, drop_last=True, num_workers=4)
test_dataloader = DataLoader(test_data, batch_size=1, shuffle=True, drop_last=True, num_workers=4)
contador = 0
y_real = []
y_pred = []
for batch, (X, y) in enumerate(train_dataloader):
    X, y = X.to(device), y.to(device)
    
    with torch.cuda.amp.autocast(): #Para evitar un desbordamiento numérico en el cálculo de la entropía cruzada
        pred = nuevo_modelo(X)
        
    _, pred_max = torch.max(pred, 1)
    
    y_real.append(y.cpu().numpy()[0])
    y_pred.append(pred_max.cpu().numpy()[0])
    
    if y != pred_max:
        contador +=1
        torchvision.utils.save_image(X.cpu() * 255, f'Mal Efficient/{batch}.png')



# liberamos memoria de GPU
import gc
gc.collect()
torch.cuda.empty_cache()
with torch.no_grad():
    torch.cuda.empty_cache()



from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, accuracy_score
clases = ['Cruz', 'Estrella', 'I', 'Linea', 'Ovalo', 'Pentagono', 'Rectangulo', 'Triangulo']
CM = confusion_matrix(y_real, y_pred)
disp = ConfusionMatrixDisplay(confusion_matrix=CM,display_labels=clases)
disp.plot(cmap=plt.cm.Oranges)
disp.ax_.set_title("Matriz de confusión GoogLeNet (Entrenamiento)")
plt.show()


