Forward pass of the UNet1DConditional model. Args: x (torch.Tensor): shape (batch_size, in_channels, time) mask (_type_): shape (batch_size, 1, time) t (_type_): shape (batch_size) spks (_type_, optional): shape: (batch_size, condition_channel
(self, x, mask, mu, t, spks=None, cond=None)
| 200 | nn.init.constant_(m.bias, 0) |
| 201 | |
| 202 | def forward(self, x, mask, mu, t, spks=None, cond=None): |
| 203 | """Forward pass of the UNet1DConditional model. |
| 204 | |
| 205 | Args: |
| 206 | x (torch.Tensor): shape (batch_size, in_channels, time) |
| 207 | mask (_type_): shape (batch_size, 1, time) |
| 208 | t (_type_): shape (batch_size) |
| 209 | spks (_type_, optional): shape: (batch_size, condition_channels). Defaults to None. |
| 210 | cond (_type_, optional): placeholder for future use. Defaults to None. |
| 211 | |
| 212 | Raises: |
| 213 | ValueError: _description_ |
| 214 | ValueError: _description_ |
| 215 | |
| 216 | Returns: |
| 217 | _type_: _description_ |
| 218 | """ |
| 219 | |
| 220 | t = self.time_embeddings(t).to(t.dtype) |
| 221 | t = self.time_mlp(t) |
| 222 | x = pack([x, mu], "b * t")[0] |
| 223 | if spks is not None: |
| 224 | spks = repeat(spks, "b c -> b c t", t=x.shape[-1]) |
| 225 | x = pack([x, spks], "b * t")[0] |
| 226 | if cond is not None: |
| 227 | x = pack([x, cond], "b * t")[0] |
| 228 | hiddens = [] |
| 229 | masks = [mask] |
| 230 | for resnet, transformer_blocks, downsample in self.down_blocks: |
| 231 | mask_down = masks[-1] |
| 232 | x = resnet(x, mask_down, t) |
| 233 | x = rearrange(x, "b c t -> b t c").contiguous() |
| 234 | attn_mask = torch.matmul(mask_down.transpose(1, 2).contiguous(), mask_down) |
| 235 | for transformer_block in transformer_blocks: |
| 236 | x = transformer_block( |
| 237 | hidden_states=x, |
| 238 | attention_mask=attn_mask, |
| 239 | timestep=t, |
| 240 | ) |
| 241 | x = rearrange(x, "b t c -> b c t").contiguous() |
| 242 | hiddens.append(x) # Save hidden states for skip connections |
| 243 | x = downsample(x * mask_down) |
| 244 | masks.append(mask_down[:, :, ::2]) |
| 245 | masks = masks[:-1] |
| 246 | mask_mid = masks[-1] |
| 247 | |
| 248 | for resnet, transformer_blocks in self.mid_blocks: |
| 249 | x = resnet(x, mask_mid, t) |
| 250 | x = rearrange(x, "b c t -> b t c").contiguous() |
| 251 | attn_mask = torch.matmul(mask_mid.transpose(1, 2).contiguous(), mask_mid) |
| 252 | for transformer_block in transformer_blocks: |
| 253 | x = transformer_block( |
| 254 | hidden_states=x, |
| 255 | attention_mask=attn_mask, |
| 256 | timestep=t, |
| 257 | ) |
| 258 | x = rearrange(x, "b t c -> b c t").contiguous() |
| 259 |
nothing calls this directly
no test coverage detected