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 musicfm.models.selector import get_model
from musicfm.data_loaders.selector import get_dataset
from musicfm.modules.lightning_module import MusicFMLitModule


def main(cfg):
    # MusicFM model
    model = get_model(cfg.model.type)(
        num_codebooks=cfg.model.num_codebooks,
        codebook_dim=cfg.model.codebook_dim,
        codebook_size=cfg.model.codebook_size,
        features=cfg.model.features,
        hop_length=cfg.model.hop_length,
        n_mels=cfg.model.n_mels,
        conv_dim=cfg.model.conv_dim,
        encoder_dim=cfg.model.encoder_dim,
        encoder_depth=cfg.model.encoder_depth,
        mask_hop=cfg.model.mask_hop,
        mask_prob=cfg.model.mask_prob,
        is_flash=cfg.model.is_flash,
        stat_path=cfg.model.stat_path,
        model_path=cfg.model.model_path,
    )

    # lightning module
    lit_module = MusicFMLitModule(
        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=False,
        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=False,
        num_workers=cfg.data.num_workers,
    )

    # callbacks
    callbacks = [
        ModelCheckpoint(
            save_last=True,
            save_top_k=cfg.core.save_top_k,
            monitor="loss_overall/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)
