import json
import os
import sys
import time
from time import gmtime, strftime

import torch
import torch.distributed as dist
import torch.nn.functional as F
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler

BASE_DIR = os.path.dirname(os.path.abspath(__file__))
PROJECT_ROOT = os.path.dirname(BASE_DIR)
sys.path.append(PROJECT_ROOT)

from finetune.config import Config
from finetune.dataset import PKLDataset
from finetune.utils.training_utils import (
    cleanup_ddp,
    format_time,
    get_model_size,
    set_seed,
    setup_ddp,
)
from model.kronos import KronosTokenizer


def create_dataloaders(config: dict, rank: int, world_size: int):
    print(f"[Rank {rank}] Creating distributed dataloaders...")
    train_dataset = PKLDataset('train')
    valid_dataset = PKLDataset('val')
    print(
        f"[Rank {rank}] Train dataset size: {len(train_dataset)}, "
        f"Validation dataset size: {len(valid_dataset)}"
    )

    train_sampler = DistributedSampler(
        train_dataset, num_replicas=world_size, rank=rank, shuffle=False
    )
    val_sampler = DistributedSampler(
        valid_dataset, num_replicas=world_size, rank=rank, shuffle=False
    )

    train_loader = DataLoader(
        train_dataset,
        batch_size=config['tokenizer_batch_size'],
        sampler=train_sampler,
        num_workers=config['num_workers'],
        pin_memory=True,
        drop_last=True,
    )
    val_loader = DataLoader(
        valid_dataset,
        batch_size=config['tokenizer_batch_size'],
        sampler=val_sampler,
        num_workers=config['num_workers'],
        pin_memory=True,
        drop_last=False,
    )
    return train_loader, val_loader, train_dataset, valid_dataset


def train_model(model, device, config, save_dir, rank, world_size):
    start_time = time.time()
    if rank == 0:
        effective_bs = config['tokenizer_batch_size'] * world_size
        print(f"[Rank {rank}] BATCHSIZE per GPU: {config['tokenizer_batch_size']}")
        print(f"[Rank {rank}] Effective total batch size: {effective_bs}")

    train_loader, val_loader, train_dataset, valid_dataset = create_dataloaders(
        config, rank, world_size
    )

    optimizer = torch.optim.AdamW(
        model.parameters(),
        lr=config['tokenizer_learning_rate'],
        weight_decay=config['adam_weight_decay'],
    )
    scheduler = torch.optim.lr_scheduler.OneCycleLR(
        optimizer=optimizer,
        max_lr=config['tokenizer_learning_rate'],
        steps_per_epoch=len(train_loader),
        epochs=config['epochs'],
        pct_start=0.03,
        div_factor=10,
    )

    best_val_loss = float('inf')
    batch_idx_global = 0
    dt_result = {}

    for epoch_idx in range(config['epochs']):
        epoch_start_time = time.time()
        model.train()
        train_loader.sampler.set_epoch(epoch_idx)

        train_dataset.set_epoch_seed(epoch_idx * 10000 + rank)
        valid_dataset.set_epoch_seed(0)

        for step_idx, (batch_x, _) in enumerate(train_loader):
            batch_x = batch_x.to(device, non_blocking=True)

            zs, bsq_loss, _, _ = model(batch_x)
            z_pre, z = zs

            recon_loss_pre = F.mse_loss(z_pre, batch_x)
            recon_loss_all = F.mse_loss(z, batch_x)
            recon_loss = recon_loss_pre + recon_loss_all
            loss = (recon_loss + bsq_loss) / 2

            optimizer.zero_grad()
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=2.0)
            optimizer.step()
            scheduler.step()

            if rank == 0 and (batch_idx_global + 1) % config['log_interval'] == 0:
                print(
                    f"[Rank {rank}, Epoch {epoch_idx + 1}/{config['epochs']}, "
                    f"Step {step_idx + 1}/{len(train_loader)}] "
                    f"LR {optimizer.param_groups[0]['lr']:.6f}, Loss: {loss.item():.4f}"
                )

            batch_idx_global += 1

        model.eval()
        tot_val_loss_sum_rank = 0.0
        val_sample_count_rank = 0
        with torch.no_grad():
            for batch_x, _ in val_loader:
                batch_x = batch_x.to(device, non_blocking=True)
                zs, bsq_loss, _, _ = model(batch_x)
                z_pre, z = zs
                recon_loss_pre = F.mse_loss(z_pre, batch_x)
                recon_loss_all = F.mse_loss(z, batch_x)
                recon_loss = recon_loss_pre + recon_loss_all
                val_loss = (recon_loss + bsq_loss) / 2

                tot_val_loss_sum_rank += val_loss.item() * batch_x.size(0)
                val_sample_count_rank += batch_x.size(0)

        val_loss_sum_tensor = torch.tensor(tot_val_loss_sum_rank, device=device)
        val_count_tensor = torch.tensor(val_sample_count_rank, device=device)
        dist.all_reduce(val_loss_sum_tensor, op=dist.ReduceOp.SUM)
        dist.all_reduce(val_count_tensor, op=dist.ReduceOp.SUM)

        avg_val_loss = (
            val_loss_sum_tensor.item() / val_count_tensor.item()
            if val_count_tensor.item() > 0
            else 0.0
        )

        if rank == 0:
            epoch_duration_sec = time.time() - epoch_start_time
            total_duration_sec = time.time() - start_time
            print(f"\n--- Epoch {epoch_idx + 1}/{config['epochs']} Summary ---")
            print(f"Validation Loss: {avg_val_loss:.4f}")
            print(f"Time This Epoch: {format_time(epoch_duration_sec)}")
            print(f"Total Time Elapsed: {format_time(total_duration_sec)}\n")

            if avg_val_loss < best_val_loss:
                best_val_loss = avg_val_loss
                save_path = f"{save_dir}/checkpoints/best_model"
                model.module.save_pretrained(save_path)
                print(f"Best model saved to {save_path} (Val Loss: {best_val_loss:.4f})")

        dist.barrier()

    dt_result['best_val_loss'] = best_val_loss
    return dt_result


def main(config: dict):
    rank, world_size, local_rank = setup_ddp()
    device = torch.device(f"cuda:{local_rank}")
    set_seed(config['seed'], rank)

    save_dir = config['tokenizer_save_dir']

    if rank == 0:
        os.makedirs(os.path.join(save_dir, 'checkpoints'), exist_ok=True)
    dist.barrier()

    train_dataset_for_dim = PKLDataset('train')
    feature_dim = len(train_dataset_for_dim.feature_list)

    model = KronosTokenizer(
        d_in=feature_dim,
        d_model=256,
        n_heads=4,
        ff_dim=512,
        n_enc_layers=4,
        n_dec_layers=4,
        ffn_dropout_p=0.0,
        attn_dropout_p=0.0,
        resid_dropout_p=0.0,
        s1_bits=10,
        s2_bits=10,
        beta=0.05,
        gamma0=1.0,
        gamma=1.1,
        zeta=0.05,
        group_size=4,
    )
    model.to(device)
    model = DDP(model, device_ids=[local_rank], find_unused_parameters=False)

    if rank == 0:
        print(f"Tokenizer Model Size: {get_model_size(model.module)}")

    dt_result = train_model(model, device, config, save_dir, rank, world_size)

    if rank == 0:
        summary = {
            'start_time': strftime("%Y-%m-%dT%H-%M-%S", gmtime()),
            'save_directory': save_dir,
            'world_size': world_size,
            'final_result': dt_result,
        }
        with open(os.path.join(save_dir, 'summary.json'), 'w', encoding='utf-8') as f:
            json.dump(summary, f, indent=4)
        print('Training finished. Summary file saved.')

    cleanup_ddp()


if __name__ == '__main__':
    # Usage: torchrun --standalone --nproc_per_node=1 finetune/train_tokenizer.py
    if "WORLD_SIZE" not in os.environ:
        raise RuntimeError("This script must be launched with `torchrun`.")

    config_instance = Config()
    config_dict = config_instance.__dict__.copy()

    main(config_dict)
