meng2003/music2dance
0
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 