File size: 3,000 Bytes
259eeac | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 | import torch.nn as nn
from torch.nn.utils import weight_norm
from .activations import SnakeBeta
from .alias_free_torch import Activation1d
def WNConv1d(*args, **kwargs):
return weight_norm(nn.Conv1d(*args, **kwargs))
class ResidualUnit(nn.Module):
def __init__(self, dim: int = 16, dilation: int = 1):
super().__init__()
pad = ((7 - 1) * dilation) // 2
self.block = nn.Sequential(
Activation1d(activation=SnakeBeta(dim, alpha_logscale=True)),
WNConv1d(dim, dim, kernel_size=7, dilation=dilation, padding=pad),
Activation1d(activation=SnakeBeta(dim, alpha_logscale=True)),
WNConv1d(dim, dim, kernel_size=1),
)
def forward(self, x):
return x + self.block(x)
class EncoderBlock(nn.Module):
def __init__(self, dim: int = 16, stride: int = 1, dilations=(1, 3, 9)):
super().__init__()
runits = [ResidualUnit(dim // 2, dilation=d) for d in dilations]
self.block = nn.Sequential(
*runits,
Activation1d(activation=SnakeBeta(dim // 2, alpha_logscale=True)),
WNConv1d(
dim // 2,
dim,
kernel_size=2 * stride,
stride=stride,
padding=stride // 2 + stride % 2,
),
)
def forward(self, x):
return self.block(x)
class SemanticEncoder(nn.Module):
def __init__(
self,
input_channels: int,
code_dim: int,
encode_channels: int,
kernel_size: int = 3,
bias: bool = True,
):
super(SemanticEncoder, self).__init__()
self.initial_conv = nn.Conv1d(
in_channels=input_channels,
out_channels=encode_channels,
kernel_size=kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
bias=False,
)
self.residual_blocks = nn.Sequential(
nn.ReLU(inplace=True),
nn.Conv1d(
encode_channels,
encode_channels,
kernel_size=kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
bias=bias,
),
nn.ReLU(inplace=True),
nn.Conv1d(
encode_channels,
encode_channels,
kernel_size=kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
bias=bias,
),
)
self.final_conv = nn.Conv1d(
in_channels=encode_channels,
out_channels=code_dim,
kernel_size=kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
bias=False,
)
def forward(self, x):
x = self.initial_conv(x) # (Batch, Encode_channels, Length)
x = self.residual_blocks(x) + x # 残差连接
x = self.final_conv(x) # (Batch, Code_dim, Length)
return x
|