Initial commit
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Core ROSVOT model components."""
|
||||
@@ -0,0 +1,295 @@
|
||||
from copy import deepcopy
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
from ...utils.commons.hparams import hparams
|
||||
from ...utils.commons.gpu_mem_track import MemTracker
|
||||
from ..commons.layers import Embedding
|
||||
from ..commons.conv import ResidualBlock, ConvBlocks
|
||||
from ..commons.conformer.conformer import ConformerLayers
|
||||
from .unet import Unet
|
||||
|
||||
def regulate_boundary(bd_logits, threshold, min_gap=18, ref_bd=None, ref_bd_min_gap=8, non_padding=None):
|
||||
# this doesn't preserve gradient
|
||||
device = bd_logits.device
|
||||
bd_logits = torch.sigmoid(bd_logits).data.cpu()
|
||||
# bd_logits[0] = bd_logits[-1] = 1e-5 # avoid itv invalid problem
|
||||
bd = (bd_logits > threshold).long()
|
||||
bd_res = torch.zeros_like(bd).long()
|
||||
for i in range(bd.shape[0]):
|
||||
bd_i = bd[i]
|
||||
last_bd_idx = -1
|
||||
start = -1
|
||||
for j in range(bd_i.shape[0]):
|
||||
if bd_i[j] == 1:
|
||||
if 0 <= start < j:
|
||||
continue
|
||||
elif start < 0:
|
||||
start = j
|
||||
else:
|
||||
if 0 <= start < j:
|
||||
if j - 1 > start:
|
||||
bd_idx = start + int(torch.argmax(bd_logits[i, start: j]).item())
|
||||
else:
|
||||
bd_idx = start
|
||||
if bd_idx - last_bd_idx < min_gap and last_bd_idx > 0:
|
||||
bd_idx = round((bd_idx + last_bd_idx) / 2)
|
||||
bd_res[i, last_bd_idx] = 0
|
||||
bd_res[i, bd_idx] = 1
|
||||
last_bd_idx = bd_idx
|
||||
start = -1
|
||||
|
||||
# assert ref_bd_min_gap <= min_gap // 2
|
||||
if ref_bd is not None and ref_bd_min_gap > 0:
|
||||
ref = ref_bd.data.cpu()
|
||||
for i in range(bd_res.shape[0]):
|
||||
ref_bd_i = ref[i]
|
||||
ref_bd_i_js = []
|
||||
for j in range(ref_bd_i.shape[0]):
|
||||
if ref_bd_i[j] == 1:
|
||||
ref_bd_i_js.append(j)
|
||||
seg_sum = torch.sum(bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap])
|
||||
if seg_sum == 0:
|
||||
bd_res[i, j] = 1
|
||||
elif seg_sum == 1 and bd_res[i, j] != 1:
|
||||
bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap] = \
|
||||
ref_bd_i[max(0, j - ref_bd_min_gap): j + ref_bd_min_gap]
|
||||
elif seg_sum > 1:
|
||||
for k in range(1, ref_bd_min_gap+1):
|
||||
if bd_res[i, max(0, j - k)] == 1 and ref_bd_i[max(0, j - k)] != 1:
|
||||
bd_res[i, max(0, j - k)] = 0
|
||||
break
|
||||
if bd_res[i, min(bd_res.shape[1] - 1, j + k)] == 1 and ref_bd_i[min(bd_res.shape[1] - 1, j + k)] != 1:
|
||||
bd_res[i, min(bd_res.shape[1] - 1, j + k)] = 0
|
||||
break
|
||||
bd_res[i, j] = 1
|
||||
# final check
|
||||
assert torch.sum(bd_res[i, ref_bd_i_js]) == len(ref_bd_i_js), \
|
||||
f"{torch.sum(bd_res[i, ref_bd_i_js])} {len(ref_bd_i_js)}"
|
||||
|
||||
bd_res = bd_res.to(device)
|
||||
|
||||
# force valid begin and end
|
||||
bd_res[:, 0] = 0
|
||||
if non_padding is not None:
|
||||
for i in range(bd_res.shape[0]):
|
||||
bd_res[i, sum(non_padding[i]) - 1:] = 0
|
||||
else:
|
||||
bd_res[:, -1] = 0
|
||||
|
||||
return bd_res
|
||||
|
||||
class BackboneNet(nn.Module):
|
||||
def __init__(self, hparams):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size = hparams['hidden_size']
|
||||
self.dropout = hparams.get('dropout', 0.0)
|
||||
updown_rates = [2, 2, 2]
|
||||
channel_multiples = [1, 1, 1]
|
||||
if hparams.get('updown_rates', None) is not None:
|
||||
updown_rates = [int(i) for i in hparams.get('updown_rates', None).split('-')]
|
||||
if hparams.get('channel_multiples', None) is not None:
|
||||
channel_multiples = [float(i) for i in hparams.get('channel_multiples', None).split('-')]
|
||||
assert len(updown_rates) == len(channel_multiples)
|
||||
# convs
|
||||
if hparams.get('bkb_net', 'conv') == 'conv':
|
||||
self.net = Unet(hidden_size, down_layers=len(updown_rates), mid_layers=hparams.get('bkb_layers', 12),
|
||||
up_layers=len(updown_rates), kernel_size=3, updown_rates=updown_rates,
|
||||
channel_multiples=channel_multiples, dropout=0, is_BTC=True,
|
||||
constant_channels=False, mid_net=None, use_skip_layer=hparams.get('unet_skip_layer', False))
|
||||
# conformer
|
||||
elif hparams.get('bkb_net', 'conv') == 'conformer':
|
||||
mid_net = ConformerLayers(
|
||||
hidden_size, num_layers=hparams.get('bkb_layers', 12), kernel_size=hparams.get('conformer_kernel', 9),
|
||||
dropout=self.dropout, num_heads=4)
|
||||
self.net = Unet(hidden_size, down_layers=len(updown_rates), up_layers=len(updown_rates), kernel_size=3,
|
||||
updown_rates=updown_rates, channel_multiples=channel_multiples, dropout=0,
|
||||
is_BTC=True, constant_channels=False, mid_net=mid_net,
|
||||
use_skip_layer=hparams.get('unet_skip_layer', False))
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
class PitchDecoder(nn.Module):
|
||||
def __init__(self, hparams):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size = hparams['hidden_size']
|
||||
self.dropout = hparams.get('dropout', 0.0)
|
||||
self.note_bd_out = nn.Linear(hidden_size, 1)
|
||||
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
|
||||
|
||||
# note prediction
|
||||
self.pitch_attn_num_head = hparams.get('pitch_attn_num_head', 1)
|
||||
self.multihead_dot_attn = nn.Linear(hidden_size, self.pitch_attn_num_head)
|
||||
self.post = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
|
||||
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
|
||||
post_net_kernel=3, act_type='leakyrelu')
|
||||
self.pitch_out = nn.Linear(hidden_size, hparams.get('note_num', 100) + 4)
|
||||
self.note_num = hparams.get('note_num', 100)
|
||||
self.note_start = hparams.get('note_start', 30)
|
||||
self.pitch_temperature = max(1e-7, hparams.get('note_pitch_temperature', 1.0))
|
||||
|
||||
def forward(self, feat, note_bd, train=True):
|
||||
bsz, T, _ = feat.shape
|
||||
|
||||
attn = torch.sigmoid(self.multihead_dot_attn(feat)) # [B, T, C] -> [B, T, num_head]
|
||||
attn = F.dropout(attn, self.dropout, train)
|
||||
attn_feat = feat.unsqueeze(3) * attn.unsqueeze(2) # [B, T, C, 1] x [B, T, 1, num_head] -> [B, T, C, num_head]
|
||||
attn_feat = torch.mean(attn_feat, dim=-1) # [B, T, C, num_head] -> [B, T, C]
|
||||
mel2note = torch.cumsum(note_bd, 1)
|
||||
note_length = torch.max(torch.sum(note_bd, dim=1)).item() + 1 # max length
|
||||
note_lengths = torch.sum(note_bd, dim=1) + 1 # [B]
|
||||
# print('note_length', note_length)
|
||||
|
||||
attn = torch.mean(attn, dim=-1, keepdim=True) # [B, T, num_head] -> [B, T, 1]
|
||||
denom = mel2note.new_zeros(bsz, note_length, dtype=attn.dtype).scatter_add_(
|
||||
dim=1, index=mel2note, src=attn.squeeze(-1)
|
||||
) # [B, T] -> [B, note_length] count the note frames of each note (with padding excluded)
|
||||
frame2note = mel2note.unsqueeze(-1).repeat(1, 1, self.hidden_size) # [B, T] -> [B, T, C], with padding included
|
||||
note_aggregate = frame2note.new_zeros(bsz, note_length, self.hidden_size, dtype=attn_feat.dtype).scatter_add_(
|
||||
dim=1, index=frame2note, src=attn_feat
|
||||
) # [B, T, C] -> [B, note_length, C]
|
||||
note_aggregate = note_aggregate / (denom.unsqueeze(-1) + 1e-5)
|
||||
note_aggregate = F.dropout(note_aggregate, self.dropout, train)
|
||||
note_logits = self.post(note_aggregate)
|
||||
note_logits = self.pitch_out(note_logits) / self.pitch_temperature
|
||||
# note_logits = torch.clamp(note_logits, min=-16., max=16.) # don't know need it or not
|
||||
|
||||
note_pred = torch.softmax(note_logits, dim=-1) # [B, note_length, note_num]
|
||||
note_pred = torch.argmax(note_pred, dim=-1) # [B, note_length]
|
||||
# for some reason, note idx maybe 130 (why?)
|
||||
note_pred[note_pred > self.note_num] = 0
|
||||
note_pred[note_pred < self.note_start] = 0
|
||||
|
||||
return note_lengths, note_logits, note_pred
|
||||
|
||||
class MidiExtractor(nn.Module):
|
||||
def __init__(self, hparams):
|
||||
super(MidiExtractor, self).__init__()
|
||||
self.hparams = deepcopy(hparams)
|
||||
self.hidden_size = hidden_size = hparams['hidden_size']
|
||||
self.dropout = hparams.get('dropout', 0.0)
|
||||
self.note_bd_threshold = hparams.get('note_bd_threshold', 0.5)
|
||||
self.note_bd_min_gap = round(hparams.get('note_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
|
||||
self.note_bd_ref_min_gap = round(hparams.get('note_bd_ref_min_gap', 50) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
|
||||
|
||||
self.mel_proj = nn.Conv1d(hparams['use_mel_bins'], hidden_size, kernel_size=3, padding=1)
|
||||
self.mel_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
|
||||
layers_in_block=2, c_multiple=1, dropout=self.dropout, num_layers=1,
|
||||
post_net_kernel=3, act_type='leakyrelu')
|
||||
self.use_pitch = hparams.get('use_pitch_embed', True)
|
||||
if self.use_pitch:
|
||||
self.pitch_embed = Embedding(300, hidden_size, 0, 'kaiming')
|
||||
self.uv_embed = Embedding(3, hidden_size, 0, 'kaiming')
|
||||
self.use_wbd = hparams.get('use_wbd', True)
|
||||
if self.use_wbd:
|
||||
self.word_bd_embed = Embedding(3, hidden_size, 0, 'kaiming')
|
||||
self.cond_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
|
||||
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
|
||||
post_net_kernel=3, act_type='leakyrelu')
|
||||
|
||||
# backbone
|
||||
self.net = BackboneNet(hparams)
|
||||
|
||||
# note bd prediction
|
||||
self.note_bd_out = nn.Linear(hidden_size, 1)
|
||||
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
|
||||
|
||||
# note prediction
|
||||
self.pitch_decoder = PitchDecoder(hparams)
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
def run_encoder(self, mel=None, word_bd=None, pitch=None, uv=None, non_padding=None):
|
||||
mel_embed = self.mel_proj(mel.transpose(1, 2)).transpose(1, 2)
|
||||
mel_embed = self.mel_encoder(mel_embed)
|
||||
pitch_embed = word_bd_embed = 0
|
||||
if self.use_pitch and pitch is not None and uv is not None:
|
||||
pitch_embed = self.pitch_embed(pitch) + self.uv_embed(uv) # [B, T, C]
|
||||
if self.use_wbd and word_bd is not None:
|
||||
word_bd_embed = self.word_bd_embed(word_bd)
|
||||
feat = self.cond_encoder(mel_embed + pitch_embed + word_bd_embed)
|
||||
|
||||
return feat
|
||||
|
||||
def forward(self, mel=None, word_bd=None, note_bd=None, pitch=None, uv=None, non_padding=None, train=True):
|
||||
ret = {}
|
||||
bsz, T, _ = mel.shape
|
||||
|
||||
feat = self.run_encoder(mel, word_bd, pitch, uv, non_padding)
|
||||
feat = self.net(feat) # [B, T, C]
|
||||
|
||||
# note bd prediction
|
||||
note_bd_logits = self.note_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.note_bd_temperature
|
||||
note_bd_logits = torch.clamp(note_bd_logits, min=-16., max=16.)
|
||||
ret['note_bd_logits'] = note_bd_logits # [B, T]
|
||||
if note_bd is None or not train:
|
||||
note_bd = regulate_boundary(note_bd_logits, self.note_bd_threshold, self.note_bd_min_gap,
|
||||
word_bd, self.note_bd_ref_min_gap, non_padding)
|
||||
ret['note_bd_pred'] = note_bd # [B, T]
|
||||
|
||||
# note pitch prediction
|
||||
note_lengths, note_logits, note_pred = self.pitch_decoder(feat, note_bd, train)
|
||||
ret['note_lengths'], ret['note_logits'], ret['note_pred'] = note_lengths, note_logits, note_pred
|
||||
|
||||
return ret
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.kaiming_normal_(self.pitch_decoder.multihead_dot_attn.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.note_bd_out.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.pitch_decoder.pitch_out.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
|
||||
nn.init.constant_(self.pitch_decoder.multihead_dot_attn.bias, 0.0)
|
||||
nn.init.constant_(self.note_bd_out.bias, 0.0)
|
||||
nn.init.constant_(self.pitch_decoder.pitch_out.bias, 0.0)
|
||||
|
||||
|
||||
class WordbdExtractor(MidiExtractor):
|
||||
def __init__(self, hparams):
|
||||
super().__init__(hparams)
|
||||
self.use_wbd = False
|
||||
self.word_bd_embed = None
|
||||
self.note_bd_out = self.note_bd_temperature = self.pitch_decoder = None
|
||||
|
||||
self.word_bd_threshold = hparams.get('word_bd_threshold', 0.5)
|
||||
self.word_bd_min_gap = round(
|
||||
hparams.get('word_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
|
||||
|
||||
self.word_bd_out = nn.Linear(self.hidden_size, 1)
|
||||
self.word_bd_temperature = max(1e-7, hparams.get('word_bd_temperature', 1.0))
|
||||
nn.init.kaiming_normal_(self.word_bd_out.weight, mode='fan_in')
|
||||
nn.init.constant_(self.word_bd_out.bias, 0.0)
|
||||
|
||||
def forward(self, mel=None, pitch=None, uv=None, non_padding=None, train=True):
|
||||
# gpu_tracker.track()
|
||||
ret = {}
|
||||
bsz, T, _ = mel.shape
|
||||
|
||||
feat = self.run_encoder(mel=mel, pitch=pitch, uv=uv, non_padding=non_padding)
|
||||
feat = self.net(feat) # [B, T, C]
|
||||
|
||||
word_bd_logits = self.word_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.word_bd_temperature
|
||||
word_bd_logits = torch.clamp(word_bd_logits, min=-16., max=16.)
|
||||
ret['word_bd_logits'] = word_bd_logits # [B, T]
|
||||
|
||||
if not train:
|
||||
word_bd = regulate_boundary(word_bd_logits, self.word_bd_threshold, self.word_bd_min_gap,
|
||||
non_padding=non_padding)
|
||||
ret['word_bd_pred'] = word_bd # [B, T]
|
||||
|
||||
return ret
|
||||
|
||||
def reset_parameters(self):
|
||||
if self.use_pitch:
|
||||
nn.init.kaiming_normal_(self.pitch_embed.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.uv_embed.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
|
||||
if self.use_pitch:
|
||||
nn.init.constant_(self.pitch_embed.weight[self.pitch_embed.padding_idx], 0.0)
|
||||
nn.init.constant_(self.uv_embed.weight[self.uv_embed.padding_idx], 0.0)
|
||||
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..commons.layers import LayerNorm, Embedding
|
||||
from ..commons.conv import ConvBlocks, ResidualBlock, get_norm_builder, get_act_builder
|
||||
|
||||
class UnetDown(nn.Module):
|
||||
def __init__(self, hidden_size, n_layers, kernel_size, down_rates, channel_multiples=None, dropout=0.0,
|
||||
is_BTC=True, constant_channels=False):
|
||||
super(UnetDown, self).__init__()
|
||||
assert n_layers == len(down_rates) # downs, down sample rate
|
||||
down_rates = [int(i) for i in down_rates]
|
||||
self.n_layers = n_layers
|
||||
self.hidden_size = hidden_size
|
||||
self.is_BTC = is_BTC
|
||||
channel_multiples = channel_multiples if channel_multiples is not None else down_rates
|
||||
self.layers = nn.ModuleList()
|
||||
self.downs = nn.ModuleList()
|
||||
in_channels = hidden_size
|
||||
for i in range(self.n_layers):
|
||||
out_channels = int(in_channels * channel_multiples[i]) if not constant_channels else in_channels
|
||||
self.layers.append(nn.Sequential(
|
||||
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
|
||||
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
|
||||
nn.Conv1d(in_channels, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
|
||||
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
|
||||
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
|
||||
))
|
||||
self.downs.append(nn.Sequential(
|
||||
nn.AvgPool1d(down_rates[i])
|
||||
))
|
||||
in_channels = out_channels
|
||||
self.last_norm = get_norm_builder('ln', out_channels)()
|
||||
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
|
||||
padding=kernel_size // 2)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
# x [B, T, C]
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2) # [B, C, T]
|
||||
skip_xs = []
|
||||
for i in range(self.n_layers):
|
||||
skip_x = self.layers[i](x)
|
||||
x = self.downs[i](skip_x)
|
||||
if self.is_BTC:
|
||||
skip_xs.append(skip_x.transpose(1, 2)) # [B, T, C]
|
||||
else:
|
||||
skip_xs.append(skip_x)
|
||||
x = self.post_net(self.last_norm(x))
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
return x, skip_xs
|
||||
|
||||
class UnetMid(nn.Module):
|
||||
def __init__(self, hidden_size, kernel_size, n_layers=None, in_dims=None, out_dims=None,
|
||||
dropout=0.0, is_BTC=True, net=None):
|
||||
super(UnetMid, self).__init__()
|
||||
in_dims = in_dims if in_dims is not None else hidden_size
|
||||
out_dims = out_dims if out_dims is not None else hidden_size
|
||||
self.pre = nn.Conv1d(in_dims, hidden_size, kernel_size, padding=kernel_size // 2)
|
||||
self.post = nn.Conv1d(hidden_size, out_dims, kernel_size, padding=kernel_size // 2)
|
||||
self.is_BTC = is_BTC
|
||||
if net is not None:
|
||||
self.net = net
|
||||
else:
|
||||
self.net = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=kernel_size,
|
||||
layers_in_block=2, c_multiple=2, dropout=dropout, num_layers=n_layers,
|
||||
post_net_kernel=3, act_type='leakyrelu', is_BTC=is_BTC)
|
||||
|
||||
def forward(self, x, cond=None, **kwargs):
|
||||
# x [B, T, C]
|
||||
if self.is_BTC:
|
||||
x = self.pre(x.transpose(1, 2)).transpose(1, 2)
|
||||
else:
|
||||
x = self.pre(x)
|
||||
if cond is None:
|
||||
cond = 0
|
||||
x = self.net(x + cond)
|
||||
if self.is_BTC:
|
||||
x = self.post(x.transpose(1, 2)).transpose(1, 2)
|
||||
else:
|
||||
x = self.post(x)
|
||||
return x
|
||||
|
||||
class UnetUp(nn.Module):
|
||||
def __init__(self, hidden_size, n_layers, kernel_size, up_rates, channel_multiples=None, dropout=0.0,
|
||||
is_BTC=True, constant_channels=False, use_skip_layer=False, skip_scale=1.0):
|
||||
super(UnetUp, self).__init__()
|
||||
assert n_layers == len(up_rates) # this is reversed in up module, from the output to the interface with middle
|
||||
up_rates = [int(i) for i in up_rates]
|
||||
self.n_layers = n_layers
|
||||
self.hidden_size = hidden_size
|
||||
self.is_BTC = is_BTC
|
||||
self.skip_scale = skip_scale
|
||||
channel_multiples = channel_multiples if channel_multiples is not None else up_rates
|
||||
# in_channels = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
|
||||
self.in_channels_lst = (np.cumprod([1] + channel_multiples) * hidden_size).astype(int) if not constant_channels \
|
||||
else [hidden_size for _ in range(self.n_layers + 1)]
|
||||
in_channels = self.in_channels_lst[-1]
|
||||
self.ups = nn.ModuleList()
|
||||
self.skip_layers = nn.ModuleList()
|
||||
self.layers = nn.ModuleList()
|
||||
for i in range(self.n_layers-1, -1, -1):
|
||||
out_channels = self.in_channels_lst[i] if not constant_channels else in_channels
|
||||
self.ups.append(nn.Sequential(
|
||||
nn.ConvTranspose1d(in_channels, in_channels, kernel_size=kernel_size, stride=up_rates[i],
|
||||
padding=kernel_size//2, output_padding=up_rates[i]-1),
|
||||
get_norm_builder('ln', in_channels)(),
|
||||
get_act_builder('leakyrelu')()
|
||||
))
|
||||
self.layers.append(nn.Sequential(
|
||||
# ResidualBlock(in_channels*2, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
|
||||
# c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
|
||||
nn.Conv1d(in_channels*2, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
|
||||
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
|
||||
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
|
||||
))
|
||||
if use_skip_layer:
|
||||
self.skip_layers.append(
|
||||
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
|
||||
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
|
||||
)
|
||||
else:
|
||||
self.skip_layers.append(nn.Identity())
|
||||
|
||||
in_channels = out_channels
|
||||
self.out_channels = out_channels
|
||||
self.last_norm = get_norm_builder('ln', out_channels)()
|
||||
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
|
||||
padding=kernel_size // 2)
|
||||
|
||||
def forward(self, x, skips, **kwargs):
|
||||
# x [B, T, C]
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2) # [B, C, T]
|
||||
for i in range(self.n_layers):
|
||||
x = self.ups[i](x)
|
||||
skip_x = skips[self.n_layers - i - 1] if not self.is_BTC \
|
||||
else skips[self.n_layers - i - 1].transpose(1, 2) # [B, T, C] -> [B, C, T]
|
||||
skip_x = self.skip_layers[i](skip_x) * self.skip_scale
|
||||
x = torch.cat((x, skip_x), dim=1) # [B, C, T]
|
||||
x = self.layers[i](x)
|
||||
x = self.post_net(self.last_norm(x))
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
return x
|
||||
|
||||
class Unet(nn.Module):
|
||||
def __init__(self, hidden_size, down_layers, up_layers, kernel_size,
|
||||
updown_rates, mid_layers=None, channel_multiples=None, dropout=0.0,
|
||||
is_BTC=True, constant_channels=False, mid_net=None, use_skip_layer=False, skip_scale=1.0):
|
||||
super(Unet, self).__init__()
|
||||
assert len(updown_rates) == down_layers == up_layers, f"{len(updown_rates)}, {down_layers}, {up_layers}"
|
||||
if channel_multiples is not None:
|
||||
assert len(channel_multiples) == len(updown_rates)
|
||||
else:
|
||||
channel_multiples = updown_rates
|
||||
self.down = UnetDown(hidden_size, down_layers, kernel_size, updown_rates,
|
||||
channel_multiples, dropout, is_BTC, constant_channels)
|
||||
down_out_dims = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
|
||||
self.mid = UnetMid(hidden_size, kernel_size, mid_layers,
|
||||
in_dims=down_out_dims, out_dims=down_out_dims, dropout=dropout, is_BTC=is_BTC, net=mid_net)
|
||||
self.up = UnetUp(hidden_size, up_layers, kernel_size, updown_rates,
|
||||
channel_multiples, dropout, is_BTC, constant_channels, use_skip_layer, skip_scale)
|
||||
|
||||
def forward(self, x, mid_cond=None, **kwargs):
|
||||
x, skips = self.down(x)
|
||||
x = self.mid(x, mid_cond)
|
||||
x = self.up(x, skips)
|
||||
return x
|
||||
Reference in New Issue
Block a user