import os
import sys
import time
import datetime
import random
import torch
import torch.nn as nn
import torch.distributed as dist
import torch.multiprocessing as mp
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler
import numpy as np
from sklearn import metrics as skm
from einops import rearrange

from ditto_v2.data_loaders.multi import MultiTaskDataset, MultiTaskBatchSampler
from ditto_v2.models.ditto import Ditto

sys.path.append("/home/minz/neon/sunoGPT/")
from torch.distributed import destroy_process_group, init_process_group
from utils.helpers import (
    dist_barrier,
    load_checkpoint,
    load_old_state_dict,
    load_old_optimizer_state_dict,
    print_with_time,
    print_with_time_master,
    save_checkpoint,
    save_old_checkpoint,
    suppress_logging,
    verify_preload_model_args,
)


ddp = int(os.environ.get("RANK", -1)) != -1  # is this a ddp run?


if ddp:
    init_process_group(backend="nccl", timeout=datetime.timedelta(seconds=24 * 60 * 60))
    ddp_rank = int(os.environ["RANK"])  # global gpu rank
    ddp_local_rank = int(os.environ["LOCAL_RANK"])  # gpu rank within node
    world_size = torch.distributed.get_world_size()  # total number of gpus
    device = f"cuda:{ddp_local_rank}"
    torch.cuda.set_device(device)
    master_process = ddp_rank == 0  # this process will do logging, checkpointing etc.
    seed_offset = ddp_rank + 1  # each process gets a different seed
    print_with_time(f"ddp init, rank {ddp_rank}, local_rank {ddp_local_rank}")
else:
    ddp_rank = 0
    ddp_local_rank = 0
    world_size = 1
    # if not ddp, we are running on a single gpu, and one process
    master_process = True
    seed_offset = 1
n_gpus_per_node = torch.cuda.device_count()
dist_barrier()
print_with_time_master(f"ddp init: world size {world_size} ddp_rank {ddp_rank}.")


def setup(rank, world_size):
    os.environ["MASTER_ADDR"] = "localhost"
    os.environ["MASTER_PORT"] = "12355"
    init_process_group("nccl", rank=rank, world_size=world_size)


def cleanup():
    destroy_process_group()


def train(rank, world_size):
    setup(rank, world_size)

    # Set device for this process
    device = torch.device(f"cuda:{rank}")
    torch.cuda.set_device(device)

    # Create model and move it to GPU with DDP
    model = Ditto(
        latent_dim=128, model_path=None, is_flash=True
    )  # Initialize your Ditto model here
    model = model.to(device)
    model = DDP(model, device_ids=[rank])

    # Define loss function and optimizer
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=learning_rate)

    # Create dataset and dataloader
    train_dataset = MultiTaskDataset(split="train")
    valid_dataset = MultiTaskDataset(split="valid")
    train_dataloader = data.DataLoader(
        dataset=train_dataset,
        batch_sampler=MultiTaskBatchSampler(train_datset, batch_size=16),
    )
    sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank)
    dataloader = DataLoader(dataset, batch_size=batch_size, sampler=sampler)

    # Training loop
    for epoch in range(num_epochs):
        model.train()
        sampler.set_epoch(epoch)
        for batch in dataloader:
            # Move batch to device
            batch = {k: v.to(device) for k, v in batch.items()}

            # Forward pass
            outputs = model(batch)
            loss = criterion(outputs, batch["labels"])

            # Backward pass and optimize
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

        # Print progress, save checkpoints, etc.
        if rank == 0:
            print(f"Epoch {epoch+1}/{num_epochs}, Loss: {loss.item()}")
            # Save checkpoint logic here

    cleanup()


if __name__ == "__main__":
    world_size = torch.cuda.device_count()
    mp.spawn(train, args=(world_size,), nprocs=world_size, join=True)


# def print_model_params(model):
#     total_params = 0
#     trainable_params = 0
#     for _, parameter in model.named_parameters():
#         params = parameter.numel()
#         total_params += params
#         if parameter.requires_grad:
#             trainable_params += params
#     print(f"Total Params: {total_params:,}, Trainable_params: {trainable_params:,}")


# class DittoModule(nn.Module):
#     def __init__(
#         self,
#         model,
#         learning_rate=1e-4,
#         gradient_clip_val=1.0,
#     ):
#         super().__init__()
#         self.lr = learning_rate
#         self.model = model
#         print_model_params(model)
#         self.loss_function = nn.CrossEntropyLoss()
#         self.gradient_clip_val = gradient_clip_val

#     def get_metrics(self, z_a, z_b, logit_scale):
#         metrics = {}

#         # scaled dot product similarity
#         logits_per_a = logit_scale * z_a @ z_b.t()
#         logits_per_b = logits_per_a.t()
#         labels = torch.arange(z_a.shape[0]).long().to(z_a.device)

#         # get cross entropy loss
#         loss = (
#             self.loss_function(logits_per_a, labels)
#             + self.loss_function(logits_per_b, labels)
#         ) / 2
#         metrics["loss"] = loss

#         # get accuracy
#         metrics["acc_a"] = skm.accuracy_score(
#             labels.detach().cpu().numpy(),
#             logits_per_a.argmax(dim=1).detach().cpu().numpy(),
#         )
#         metrics["acc_b"] = skm.accuracy_score(
#             labels.detach().cpu().numpy(),
#             logits_per_b.argmax(dim=1).detach().cpu().numpy(),
#         )

#         # get ranking metrics
#         logits = {
#             "a_to_b": logits_per_a.detach().cpu(),
#             "b_to_a": logits_per_b.detach().cpu(),
#         }
#         ground_truth = torch.arange(len(z_b)).view(-1, 1)
#         for name, logit in logits.items():
#             ranking = torch.argsort(logit, descending=True)
#             preds = torch.where(ranking == ground_truth)[1]
#             preds = preds.detach().cpu().numpy()
#             metrics[f"{name}_mean_rank"] = preds.mean() + 1
#             metrics[f"{name}_median_rank"] = np.floor(np.median(preds)) + 1
#             for k in [1, 5, 10]:
#                 metrics[f"{name}_R@{k}"] = np.mean(preds < k)
#             metrics[f"{name}_mAP@10"] = np.mean(
#                 np.where(preds < 10, 1 / (preds + 1), 0.0)
#             )
#         return metrics

#     def forward(self, batch, stage):
#         # get batch
#         inps, task = batch
#         task = task[0]
#         inp1, inp2 = inps

#         # forward based on task
#         if task in ["self_sim", "self_vox_sim", "artist_sim", "artist_vox_sim", "album_sim"]:
#             if stage == "train":
#                 random_length1 = random.randint(24000 * 5, 24000 * 15)
#                 random_length2 = random.randint(24000 * 5, 24000 * 15)
#                 inp1 = inp1[:, :random_length1]
#                 inp2 = inp2[:, :random_length2]
#             outputs = self.model(inp1, inp2, task)
#             modality_ab = "m2m"
#             modality_ba = "m2m"
#         elif task in ["self_lyric_sim"]:
#             if stage == "train":
#                 random_length1 = random.randint(50, 313)
#                 random_length2 = random.randint(50, 313)
#                 inp1 = [line[:random_length1] for line in inp1]
#                 inp2 = [line[:random_length2] for line in inp2]
#             outputs = self.model(inp1, inp2, task)
#             modality_ab = "t2t"
#             modality_ba = "t2t"
#         elif task in ["genre_sim", "lyric_sim"]:
#             if stage == "train":
#                 random_length1 = random.randint(24000 * 5, 24000 * 15)
#                 random_length2 = random.randint(50, 313)
#                 inp1 = inp1[:, :random_length1]
#                 inp2 = [line[:random_length2] for line in inp2]
#             outputs = self.model(inp1, inp2, task)
#             modality_ab = "m2t"
#             modality_ba = "t2m"

#         inp1_emb, inp2_emb, logit_scale = outputs

#         # get metrics
#         metrics = self.get_metrics(inp1_emb, inp2_emb, logit_scale)

#         return metrics, modality_ab, modality_ba

# def train(rank, world_size, cfg):
#     dist.init_process_group("nccl", rank=rank, world_size=world_size)
#     torch.cuda.set_device(rank)

#     model = Ditto(
#         latent_dim=cfg.model.latent_dim,
#         model_path=cfg.model.model_path,
#         is_flash=True
#     )
#     model = model.to(rank)

#     ditto_module = DittoModule(
#         model=model,
#         learning_rate=cfg.optim.learning_rate,
#     )
#     ditto_module = DDP(ditto_module, device_ids=[rank])

#     optimizer = torch.optim.AdamW(
#         [
#             {"params": ditto_module.module.model.music_projection.parameters(), "lr": ditto_module.module.lr},
#             {"params": ditto_module.module.model.text_projection.parameters(), "lr": ditto_module.module.lr},
#             {"params": ditto_module.module.model.music_encoder.parameters(), "lr": ditto_module.module.lr / 10},
#             {"params": ditto_module.module.model.text_encoder.parameters(), "lr": ditto_module.module.lr / 10},
#         ],
#         lr=ditto_module.module.lr,
#     )

#     train_dataset = MultiTaskDataset(split="train")
#     valid_dataset = MultiTaskDataset(split="valid")

#     train_sampler = DistributedSampler(train_dataset, num_replicas=world_size, rank=rank)
#     valid_sampler = DistributedSampler(valid_dataset, num_replicas=world_size, rank=rank)

#     train_dataloader = DataLoader(
#         dataset=train_dataset,
#         batch_sampler=MultiTaskBatchSampler(train_dataset, batch_size=cfg.data.batch_size),
#         num_workers=cfg.data.num_workers,
#         sampler=train_sampler,
#     )
#     validation_dataloader = DataLoader(
#         dataset=valid_dataset,
#         batch_sampler=MultiTaskBatchSampler(valid_dataset, batch_size=cfg.data.batch_size),
#         num_workers=cfg.data.num_workers,
#         sampler=valid_sampler,
#     )

#     for epoch in range(cfg.core.max_epochs):
#         ditto_module.train()
#         for batch in train_dataloader:
#             optimizer.zero_grad()
#             metrics, modality_ab, modality_ba = ditto_module(batch, "train")
#             loss = metrics["loss"]
#             loss.backward()
#             torch.nn.utils.clip_grad_norm_(ditto_module.parameters(), ditto_module.module.gradient_clip_val)
#             optimizer.step()

#             if rank == 0:
#                 print(f"Epoch {epoch}, Loss: {loss.item()}")
#                 print(f"mAP10-{modality_ab}_train: {metrics['a_to_b_mAP@10']}")
#                 print(f"mAP10-{modality_ba}_train: {metrics['b_to_a_mAP@10']}")

#         ditto_module.eval()
#         val_loss = 0
#         with torch.no_grad():
#             for batch in validation_dataloader:
#                 metrics, modality_ab, modality_ba = ditto_module(batch, "validation")
#                 val_loss += metrics["loss"].item()

#         val_loss /= len(validation_dataloader)
#         if rank == 0:
#             print(f"Validation Loss: {val_loss}")

# def main(cfg):
#     world_size = torch.cuda.device_count()
#     mp.spawn(train, args=(world_size, cfg), nprocs=world_size, join=True)

#     os.environ['MASTER_ADDR'] = 'localhost'
#     os.environ['MASTER_PORT'] = '12355'

#     mp.spawn(train, args=(world_size, cfg), nprocs=world_size, join=True)


# if __name__ == "__main__":
#     import sys
#     import warnings
#     warnings.filterwarnings("ignore", category=FutureWarning)
#     import omegaconf

#     cfg = omegaconf.OmegaConf.load(sys.argv[1])
#     main(cfg)
