import os

os.environ["TOKENIZERS_PARALLELISM"] = "false"
import sys
import warnings

warnings.filterwarnings("ignore", category=FutureWarning)
import omegaconf
import lightning as L
from torch.utils import data
from lightning.pytorch.callbacks import ModelCheckpoint
from lightning.pytorch.loggers import WandbLogger

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


def main(cfg):
    model = Ditto(
        latent_dim=cfg.model.latent_dim, model_path=cfg.model.model_path, is_flash=True
    )

    # lightning module
    lit_module = DittoLitModule(
        model=model,
        learning_rate=cfg.optim.learning_rate,
    )

    # data loaders
    train_dataset = MultiTaskDataset(split="train")
    valid_dataset = MultiTaskDataset(split="valid")
    train_dataloader = data.DataLoader(
        dataset=train_dataset,
        batch_sampler=MultiTaskBatchSampler(
            train_dataset, batch_size=cfg.data.batch_size
        ),
        num_workers=cfg.data.num_workers,
    )
    validation_dataloader = data.DataLoader(
        dataset=valid_dataset,
        batch_sampler=MultiTaskBatchSampler(
            valid_dataset, batch_size=cfg.data.batch_size
        ),
        num_workers=cfg.data.num_workers,
    )

    # callbacks
    callbacks = [
        ModelCheckpoint(
            save_last=True,
            save_top_k=cfg.core.save_top_k,
            monitor="loss_validation",
            mode="min",
            dirpath="/home/minz/logs/%s" % cfg.core.version,
        )
    ]

    # logger
    logger = WandbLogger(
        name=cfg.core.version, save_dir="/app/suno/minz/wandb_logs", log_model="all"
    )

    # trainer
    trainer = L.Trainer(
        accelerator="gpu",
        devices=cfg.core.devices,
        num_nodes=cfg.core.num_nodes,
        strategy="deepspeed",
        precision=cfg.core.precision,
        limit_train_batches=cfg.data.limit_train,
        profiler="simple",  # "simple" or "advanced"
        callbacks=callbacks,
        max_epochs=cfg.core.max_epochs,
        logger=logger,
        use_distributed_sampler=False,
    )
    trainer.fit(
        lit_module,
        train_dataloader,
        validation_dataloader,
        ckpt_path=cfg.core.ckpt_path,
    )


if __name__ == "__main__":
    cfg = omegaconf.OmegaConf.load(sys.argv[1])
    main(cfg)
