| 173 | |
| 174 | # Nonstreaming Vocos backbone based on Transformer layers |
| 175 | class VocosBackbone(nn.Module): |
| 176 | def __init__( |
| 177 | self, |
| 178 | embed_dim: int = 1024, |
| 179 | num_layers: int = 12, |
| 180 | num_heads: int = 16, |
| 181 | dropout: float = 0.1, |
| 182 | ): |
| 183 | super().__init__() |
| 184 | self.in_proj = nn.Conv1d(embed_dim, embed_dim, kernel_size=7, padding=3) |
| 185 | self.prior_net = nn.Sequential( |
| 186 | ResnetBlock(embed_dim, embed_dim, dropout=dropout), |
| 187 | ResnetBlock(embed_dim, embed_dim, dropout=dropout), |
| 188 | ) |
| 189 | self.transformers = nn.ModuleList( |
| 190 | [WhisperEncoderLayer(embed_dim, num_heads) for _ in range(num_layers)] |
| 191 | ) |
| 192 | self.post_net = nn.Sequential( |
| 193 | ResnetBlock(embed_dim, embed_dim, dropout=dropout), |
| 194 | ResnetBlock(embed_dim, embed_dim, dropout=dropout), |
| 195 | ) |
| 196 | self.final_norm = nn.LayerNorm(embed_dim, eps=1e-6) |
| 197 | |
| 198 | def forward( |
| 199 | self, |
| 200 | x: torch.Tensor, |
| 201 | x_lens: torch.Tensor, |
| 202 | ): |
| 203 | """ |
| 204 | Args: |
| 205 | x: shape (b, t, c) |
| 206 | x_lens: shape (b,) |
| 207 | """ |
| 208 | x = x.transpose(1, 2) |
| 209 | x = self.in_proj(x) |
| 210 | x = self.prior_net(x) |
| 211 | x = x.transpose(1, 2) |
| 212 | |
| 213 | attention_mask = make_nonpad_mask(x_lens).unsqueeze(1) # (b, 1, t) |
| 214 | # NOTE(sfy): I think positional embedding is unnecessary |
| 215 | for layer in self.transformers: |
| 216 | x = layer(x, attention_mask) |
| 217 | x = x.transpose(1, 2) |
| 218 | x = self.post_net(x) |
| 219 | x = x.transpose(1, 2) |
| 220 | x = self.final_norm(x) |
| 221 | return x |
| 222 | |
| 223 | |
| 224 | # Streaming Vocos backbone based on Transformer layers |