import torch
import json
from torch.optim.lr_scheduler import LambdaLR
import wandb
from data import Vocab
from data_gen import DataGenerator, Dataset, Augmenter
from train_target import train_target_wants_text
from model import (
    ModelArgs,
    MusicalPositionEmbedTransformer,
    LabelSmoothing,
    SimpleLossCompute,
    RMSNorm,
)
from main import TrainState, run_epoch, rate
from embedder import Embedder

import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
import dataclasses
import logging
import transformers.models.t5.modeling_t5 as t5
import dataset_classes
import data_aug
from pathlib import Path


def run(config_path):
    with open(config_path, "r") as f:
        config = json.load(f)

    dist.init_process_group("nccl")
    rank = dist.get_rank()
    world_size = dist.get_world_size()
    device_id = rank % torch.cuda.device_count()
    torch.cuda.set_device(device_id)

    logging.info(f"start rank {rank+1} of {world_size}")

    vocab = Vocab.from_config(config)

    train_target = str(config["train_target"])
    enable_text_prompt = train_target_wants_text(train_target)

    batch_size = int(config["batch_size"])
    seq_len = int(config["seq_len"])
    seq_len_min = int(config["seq_len_min"])
    seq_len_max = int(config["seq_len_max"])

    if enable_text_prompt:
        encoder = Embedder()
    else:
        encoder = None

    # construct dataset and generators from config
    train_split_idx = int(config["train_split_idx"])
    eval_split_idx = int(config["eval_split_idx"])
    example_continuous = bool(config["example_continuous"])
    pack_batch = bool(config["pack_batch"])

    dataset = Dataset.dynamic_from_config(config, "dataset_class")
    augmenter = Augmenter.dynamic_from_config(vocab, config, "augmenter_class")

    train_data_gen = DataGenerator(
        dataset,
        vocab,
        split=train_split_idx,
        ranksize=(rank, world_size),
        batch_size=batch_size,
        seq_len=seq_len,
        seq_len_min=seq_len_min,
        seq_len_max=seq_len_max,
        train_target=train_target,
        parallelism=8,
        to_device=f"cuda:{device_id}",
        text_tokenize=encoder.tokenize if enable_text_prompt else None,
        example_continuous=example_continuous,
        pack_batch=pack_batch,
        augmenter=augmenter,
    )
    eval_data_gen = DataGenerator(
        dataset,
        vocab,
        split=eval_split_idx,
        ranksize=(rank, world_size),
        batch_size=batch_size,
        seq_len=seq_len,
        seq_len_min=seq_len_min,
        seq_len_max=seq_len_max,
        train_target=train_target,
        parallelism=8,
        to_device=f"cuda:{device_id}",
        text_tokenize=encoder.tokenize if enable_text_prompt else None,
        example_continuous=example_continuous,
        pack_batch=False,
        augmenter=augmenter,
    )
    criterion = LabelSmoothing(
        padding_idx=0, smoothing=float(config["label_smoothing"])
    )
    model = MusicalPositionEmbedTransformer(
        vocab,
        dataclasses.replace(
            ModelArgs.from_config(vocab, config),
            cache=False,
        ),
        encoder,
    )

    if rank == 0:
        print(model)
        print(f"model size: {model.param_count:.2e} parameters")

    # create model and move it to GPU with id rank
    model = model.to(device_id)
    model = DDP(model, device_ids=[device_id])

    lr = 0.1
    lr_factor = float(config["lr_factor"])

    num_batches = torch.zeros(2, dtype=torch.long, device="cuda")
    if rank == 0:
        train_split_size = train_data_gen.num_examples()
        eval_split_size = eval_data_gen.num_examples()
        print(
            f"train split size {train_split_size:.2e}, eval split size {eval_split_size:.2e}"
        )
        num_batches[0] = train_split_size // (batch_size * world_size)
        num_batches[1] = eval_split_size // (batch_size * world_size)
        assert num_batches[0] > 0, "train split is too small"
        assert num_batches[1] > 0, "eval split is too small"

    dist.broadcast(num_batches, 0)
    num_batches, num_eval_batches = num_batches.tolist()

    accum_iter = int(config["accum_iter"])
    eval_iter = int(config["eval_iter"])
    epochs = int(config["epochs"])
    lr_decay_epochs = int(config["lr_decay_epochs"])

    # choose whether to weight decay each module
    # lifted from minGPT
    decay = set()
    no_decay = set()
    whitelist_weight_modules = (torch.nn.Linear,)
    blacklist_weight_modules = (
        torch.nn.LayerNorm,
        torch.nn.Embedding,
        RMSNorm,
        t5.T5LayerNorm,
    )
    for mn, m in model.named_modules():
        for pn, p in m.named_parameters():
            fpn = "%s.%s" % (mn, pn) if mn else pn  # full param name

            if pn.endswith("bias") or pn.endswith("alpha"):
                no_decay.add(fpn)
            elif pn.endswith("weight") and isinstance(m, whitelist_weight_modules):
                decay.add(fpn)
            elif pn.endswith("weight") and isinstance(m, blacklist_weight_modules):
                no_decay.add(fpn)

    # validate that we considered every parameter
    param_dict = {pn: p for pn, p in model.named_parameters()}
    inter_params = decay & no_decay
    union_params = decay | no_decay
    assert (
        len(inter_params) == 0
    ), "parameters %s made it into both decay/no_decay sets!" % (str(inter_params),)
    assert (
        len(param_dict.keys() - union_params) == 0
    ), "parameters %s were not separated into either decay/no_decay set!" % (
        str(param_dict.keys() - union_params),
    )

    # create the pytorch optimizer object
    optim_groups = [
        {
            "params": [param_dict[pn] for pn in sorted(list(decay))],
            "weight_decay": float(config["weight_decay"]),
        },
        {
            "params": [param_dict[pn] for pn in sorted(list(no_decay))],
            "weight_decay": 0.0,
        },
    ]

    optimizer = torch.optim.AdamW(
        optim_groups,
        lr=lr,
        betas=(0.9, 0.99),
        eps=1e-7,
    )
    lr_scheduler = LambdaLR(
        optimizer=optimizer,
        lr_lambda=lambda step: rate(
            step,
            model_size=model.module.params.dim,
            factor=lr_factor,
            min_factor=0.1,
            steps_in_epoch=num_batches * lr_decay_epochs,
        ),
    )

    max_examples_per_rank = config.get("max_examples_per_rank", None)
    enable_checkpoint = bool(config["enable_checkpoint"])
    checkpoint_path = config.get("checkpoint_path", "./checkpoints")
    enable_grad_scaler = bool(config["enable_grad_scaler"])
    enable_shuffle = bool(config["enable_shuffle"])
    enable_eval = bool(config["enable_eval"])
    enable_cross_attention = bool(config["enable_cross_attention"])

    Path(checkpoint_path).mkdir(parents=True, exist_ok=True)

    scaler = (
        torch.cuda.amp.GradScaler(growth_interval=200) if enable_grad_scaler else None
    )

    if rank == 0:
        wandb.init(
            project="composer-v18",
            config={
                "batch_size": batch_size,
                "num_batches": num_batches,
                "accum_iter": accum_iter,
                "eval_iter": eval_iter,
                "lr_factor_times_world_size": lr_factor,
                "epochs": epochs,
                "layers": model.module.params.n_layers,
                "heads": model.module.params.n_heads,
                "dim": model.module.params.dim,
                "dropout": model.module.params.dropout,
                "label_smoothing": criterion.smoothing,
                "world_size": world_size,
                "max_examples_per_rank": max_examples_per_rank,
                "param_count": model.module.param_count,
                "enable_flash": model.module.params.enable_flash,
                "enable_grad_scaler": enable_grad_scaler,
                "enable_text_prompt": enable_text_prompt,
                "enable_eval": enable_eval,
                "enable_shuffle": enable_shuffle,
                "train_target": train_target,
                "enable_cross_attention": enable_cross_attention,
            },
        )
        wandb.watch(model, log=None)

    model.train()
    train_state = TrainState()
    for epoch in range(epochs):
        if rank == 0 and enable_shuffle:
            train_data_gen.shuffle()
        dist.barrier()
        _, _, _, train_state = run_epoch(
            train_data_gen.generate(
                force_batches=num_batches,
                order="random",
                seq_len_max=seq_len_max,
            ),
            lambda: eval_data_gen.generate(
                order="id", seq_len_max=seq_len_max, force_batches=100
            ),  # XXX
            model,
            SimpleLossCompute(criterion),
            optimizer,
            lr_scheduler,
            batch_size,
            mode="train",
            num_batches=num_batches,
            num_eval_batches=num_eval_batches,
            accum_iter=accum_iter,
            eval_iter=eval_iter if enable_eval else None,
            enable_checkpoint=enable_checkpoint and rank == 0,
            epoch=epoch,
            quiet=rank != 0,
            max_examples_per_rank=max_examples_per_rank,
            scaler=scaler,
            train_state=train_state,
            world_size=world_size,
            checkpoint_path=checkpoint_path,
        )


if __name__ == "__main__":
    import sys
    from util import configure_logging

    configure_logging()

    run(sys.argv[1])
    wandb.finish()
