import random import torch import torch.nn.functional as F import numpy as np from torch import nn from ditto_v2.modules.music_encoder import MusicEncoder from ditto_v2.modules.text_encoder import TextEncoder from ditto_v2.modules.mlp import Projection class Ditto(nn.Module): def __init__( self, latent_dim=512, model_path=None, is_flash=True, ): super().__init__() self.latent_dim = latent_dim # prepare encoders self.music_encoder = MusicEncoder(layer_ix=12, is_flash=is_flash) self.text_encoder = TextEncoder() # 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) 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 # random.seed(142) # SS["music_encoder.music_encoder.cls_token"] = nn.Parameter(torch.randn(1024)) self.load_state_dict(SS, strict=True) print("model loaded!") def get_projection_layers(self): music_dim = 1024 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, task): music_emb = self.music_projection.float()(self.music_encoder(wav, task).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_multi(self, wav, text, task): # music encoding music_emb = self.music_projection.float()(self.music_encoder(wav, task).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() def forward_text(self, text1, text2): # text encoding text1_emb = self.text_projection.float()(self.text_encoder(text1).float()) text1_emb = F.normalize(text1_emb, dim=-1) # text encoding text2_emb = self.text_projection.float()(self.text_encoder(text2).float()) text2_emb = F.normalize(text2_emb, dim=-1) return text1_emb, text2_emb, self.logit_scale.exp() def forward_music(self, wav1, wav2, task): # music encoding music1_emb = self.music_projection.float()( self.music_encoder(wav1, task).float() ) music1_emb = F.normalize(music1_emb, dim=-1) # music encoding music2_emb = self.music_projection.float()( self.music_encoder(wav2, task).float() ) music2_emb = F.normalize(music2_emb, dim=-1) return music1_emb, music2_emb, self.logit_scale.exp() def forward(self, inp1, inp2, task): if task in [ "self_sim", "self_vox_sim", "artist_sim", "artist_vox_sim", "album_sim", ]: return self.forward_music(inp1, inp2, task) elif task in ["self_lyric_sim"]: return self.forward_text(inp1, inp2) elif task in ["genre_sim", "lyric_sim"]: # music first return self.forward_multi(inp1, inp2, task)