| 79 | |
| 80 | |
| 81 | class DiffNet(nn.Module): |
| 82 | def __init__(self, in_dims=80): |
| 83 | super().__init__() |
| 84 | self.params = params = AttrDict( |
| 85 | # Model params |
| 86 | encoder_hidden=hparams['hidden_size'], |
| 87 | residual_layers=hparams['residual_layers'], |
| 88 | residual_channels=hparams['residual_channels'], |
| 89 | dilation_cycle_length=hparams['dilation_cycle_length'], |
| 90 | ) |
| 91 | self.input_projection = Conv1d(in_dims, params.residual_channels, 1) |
| 92 | self.diffusion_embedding = SinusoidalPosEmb(params.residual_channels) |
| 93 | dim = params.residual_channels |
| 94 | self.mlp = nn.Sequential( |
| 95 | nn.Linear(dim, dim * 4), |
| 96 | Mish(), |
| 97 | nn.Linear(dim * 4, dim) |
| 98 | ) |
| 99 | self.residual_layers = nn.ModuleList([ |
| 100 | ResidualBlock(params.encoder_hidden, params.residual_channels, 2 ** (i % params.dilation_cycle_length)) |
| 101 | for i in range(params.residual_layers) |
| 102 | ]) |
| 103 | self.skip_projection = Conv1d(params.residual_channels, params.residual_channels, 1) |
| 104 | self.output_projection = Conv1d(params.residual_channels, in_dims, 1) |
| 105 | nn.init.zeros_(self.output_projection.weight) |
| 106 | |
| 107 | def forward(self, spec, diffusion_step, cond): |
| 108 | """ |
| 109 | |
| 110 | :param spec: [B, 1, M, T] |
| 111 | :param diffusion_step: [B, 1] |
| 112 | :param cond: [B, M, T] |
| 113 | :return: |
| 114 | """ |
| 115 | x = spec[:, 0] |
| 116 | x = self.input_projection(x) # x [B, residual_channel, T] |
| 117 | |
| 118 | x = F.relu(x) |
| 119 | diffusion_step = self.diffusion_embedding(diffusion_step) |
| 120 | diffusion_step = self.mlp(diffusion_step) |
| 121 | skip = [] |
| 122 | for layer_id, layer in enumerate(self.residual_layers): |
| 123 | x, skip_connection = layer(x, cond, diffusion_step) |
| 124 | skip.append(skip_connection) |
| 125 | |
| 126 | x = torch.sum(torch.stack(skip), dim=0) / sqrt(len(self.residual_layers)) |
| 127 | x = self.skip_projection(x) |
| 128 | x = F.relu(x) |
| 129 | x = self.output_projection(x) # [B, 80, T] |
| 130 | return x[:, None, :, :] |
no outgoing calls
no test coverage detected