import torch import torch.nn.functional as F import numpy as np from torch import nn from ditto.modules.music_encoder import MusicEncoder from ditto.modules.text_encoder import TextEncoder from ditto.modules.mlp import Projection, MLP class Ditto(nn.Module): def __init__( self, music_encoder_name="musicfm_mertlong", text_encoder_name="xlm-roberta", latent_dim=512, model_path=None, is_flash=True, ): super().__init__() self.music_encoder_name = music_encoder_name self.text_encoder_name = text_encoder_name self.latent_dim = latent_dim # prepare encoders self.music_encoder = MusicEncoder( music_encoder_name, layer_ix=12, is_flash=is_flash ) self.text_encoder = TextEncoder(text_encoder_name) # get projection layers self.music_projection, self.text_projection = self.get_projection_layers() # logit scale self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07)) # load model if model_path: S = torch.load(model_path)["state_dict"] SS = {k[6:]: v for k, v in S.items()} SS["logit_scale"] = nn.Parameter( torch.ones([]) * np.log(1 / 0.07) ) # TODO: include this scale in back propagation self.load_state_dict(SS, strict=True) print("model loaded!") def get_projection_layers(self): if self.music_encoder_name == "musicfm_mertlong": music_dim = 1024 elif self.music_encoder_name == "musicfm_concat": music_dim = 1024 if self.text_encoder_name == "xlm-roberta": text_dim = 768 music_projection = Projection(music_dim, self.latent_dim) text_projection = Projection(text_dim, self.latent_dim) return music_projection, text_projection @torch.no_grad() def music_to_latent(self, wav): music_emb = self.music_projection.float()(self.music_encoder(wav).float()) music_emb = F.normalize(music_emb, dim=-1) return music_emb @torch.no_grad() def text_to_latent(self, text): text_emb = self.text_projection.float()(self.text_encoder(text).float()) text_emb = F.normalize(text_emb, dim=-1) return text_emb def forward(self, wav, text): # music encoding music_emb = self.music_projection.float()(self.music_encoder(wav).float()) music_emb = F.normalize(music_emb, dim=-1) # text encoding text_emb = self.text_projection.float()(self.text_encoder(text).float()) text_emb = F.normalize(text_emb, dim=-1) return music_emb, text_emb, self.logit_scale.exp()