CoolFace
Apppublic

meng2003/music2dance

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
transflowpp_model.py131 linesDownload Raw Back to models
1import torch2from torch import nn3from .transformer import BasicTransformerModel4from models import BaseModel5from models.flowplusplus import FlowPlusPlus6import ast7from .util.generation import autoregressive_generation_multimodal8 9class TransFlowppModel(BaseModel):10    def __init__(self, opt):11        super().__init__(opt)12        input_mods = self.input_mods13        output_mods = self.output_mods14        dins = self.dins15        douts = self.douts16        input_lengths = self.input_lengths17        output_lengths = self.output_lengths18 19        self.input_mod_nets = []20        self.output_mod_nets = []21        self.output_mod_glows = []22        self.module_names = []23        for i, mod in enumerate(input_mods):24            net = BasicTransformerModel(opt.dhid, dins[i], opt.nhead, opt.dhid, 2, opt.dropout, self.device, use_pos_emb=True, input_length=input_lengths[i]).to(self.device)25            name = "_input_"+mod26            setattr(self,"net"+name, net)27            self.input_mod_nets.append(net)28            self.module_names.append(name)29        for i, mod in enumerate(output_mods):30            net = BasicTransformerModel(opt.dhid, opt.dhid, opt.nhead, opt.dhid, opt.nlayers, opt.dropout, self.device, use_pos_emb=True, input_length=sum(input_lengths)).to(self.device)31            name = "_output_"+mod32            setattr(self, "net"+name, net)33            self.output_mod_nets.append(net)34            self.module_names.append(name)35 36            # import pdb;pdb.set_trace()37            glow = FlowPlusPlus(scales=ast.literal_eval(opt.scales),38                                     in_shape=(douts[i], output_lengths[i], 1),39                                     cond_dim=opt.dhid,40                                     mid_channels=opt.dhid,41                                     num_blocks=opt.num_glow_coupling_blocks,42                                     num_components=opt.num_mixture_components,43                                     use_attn=opt.glow_use_attn,44                                     use_logmix=opt.num_mixture_components>0,45                                     drop_prob=opt.dropout46                                     )47            name = "_output_glow_"+mod48            setattr(self, "net"+name, glow)49            self.output_mod_glows.append(glow)50 51 52        # self.generate_full_masks()53        self.inputs = []54        self.targets = []55        self.criterion = nn.MSELoss()56 57    def name(self):58        return "Transformerflow"59 60    @staticmethod61    def modify_commandline_options(parser, opt):62        parser.add_argument('--dhid', type=int, default=512)63        parser.add_argument('--nlayers', type=int, default=6)64        parser.add_argument('--nhead', type=int, default=8)65        parser.add_argument('--dropout', type=float, default=0.1)66        parser.add_argument('--scales', type=str, default="[[10,0]]")67        parser.add_argument('--num_glow_coupling_blocks', type=int, default=10)68        parser.add_argument('--num_mixture_components', type=int, default=0)69        parser.add_argument('--glow_use_attn', action='store_true', help="whether to use the internal attention for the FlowPlusPLus model")70        return parser71 72    # def generate_full_masks(self):73    #     input_mods = self.input_mods74    #     output_mods = self.output_mods75    #     input_lengths = self.input_lengths76    #     self.src_masks = []77    #     for i, mod in enumerate(input_mods):78    #         mask = torch.zeros(input_lengths[i],input_lengths[i])79    #         self.register_buffer('src_mask_'+str(i), mask)80    #         self.src_masks.append(mask)81    #82    #     self.output_masks = []83    #     for i, mod in enumerate(output_mods):84    #         mask = torch.zeros(sum(input_lengths),sum(input_lengths))85    #         self.register_buffer('out_mask_'+str(i), mask)86    #         self.output_masks.append(mask)87 88    def forward(self, data):89        # in lightning, forward defines the prediction/inference actions90        latents = []91        for i, mod in enumerate(self.input_mods):92            # mask = getattr(self,"src_mask_"+str(i))93            #mask = self.src_masks[i]94            latents.append(self.input_mod_nets[i].forward(data[i]))95        latent = torch.cat(latents)96        outputs = []97        for i, mod in enumerate(self.output_mods):98            # mask = getattr(self,"out_mask_"+str(i))99            #mask = self.output_masks[i]100            trans_output = self.output_mod_nets[i].forward(latent)[:self.output_lengths[i]]101            output, _ = self.output_mod_glows[i](x=None, cond=trans_output.permute(1,0,2), reverse=True)102            outputs.append(output.permute(1,0,2))103 104        # import pdb;pdb.set_trace()105        #shape106 107        return outputs108 109    def training_step(self, batch, batch_idx):110        self.set_inputs(batch)111        latents = []112        for i, mod in enumerate(self.input_mods):113            # mask = getattr(self,"src_mask_"+str(i))114            latents.append(self.input_mod_nets[i].forward(self.inputs[i]))115 116        latent = torch.cat(latents)117        loss = 0118        for i, mod in enumerate(self.output_mods):119            # mask = getattr(self,"out_mask_"+str(i))120            output = self.output_mod_nets[i].forward(latent)[:self.output_lengths[i]]121            glow = self.output_mod_glows[i]122            # import pdb;pdb.set_trace()123            z, sldj = glow(x=self.targets[i].permute(1,0,2), cond=output.permute(1,0,2)) #time, batch, features -> batch, time, features124            loss += glow.loss_generative(z, sldj)125        self.log('nll_loss', loss)126        return loss127 128    #def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx,129    #                           optimizer_closure, on_tpu, using_native_amp, using_lbfgs):130    #    optimizer.zero_grad()131