import sys
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.data_loaders.selector import get_dataset
from ditto.models.ditto import Ditto
from ditto.modules.lightning_module import DittoLitModule


def main(cfg):
    model = Ditto(
        music_encoder_name=cfg.model.music_encoder_name,
        text_encoder_name=cfg.model.text_encoder_name,
        latent_dim=cfg.model.latent_dim,
        model_path=cfg.model.model_path,
    )

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

    # data loaders
    train_dataloader = data.DataLoader(
        dataset=get_dataset(
            cfg.data.train_dataset,
            split="train",
            num_samples=cfg.data.num_train_samples,
        ),
        batch_size=cfg.data.batch_size,
        shuffle=True,
        drop_last=True,
        num_workers=cfg.data.num_workers,
    )
    validation_dataloader = data.DataLoader(
        dataset=get_dataset(
            cfg.data.valid_dataset,
            split="valid",
            num_samples=cfg.data.num_val_samples,
        ),
        batch_size=cfg.data.batch_size,
        shuffle=False,
        drop_last=True,
        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,
    )
    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)
