import sys

sys.path.append("/home/minz/glockenspiel/musicfm-training/")
import omegaconf
import hydra
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.contrastive_model 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,
    )

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

    # data loaders
    test_dataloader = data.DataLoader(
        dataset=get_dataset("genius_lyric")(split="valid"),
        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.test(
        lit_module, dataloaders=[test_dataloader], ckpt_path=cfg.core.ckpt_path
    )


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