MCPcopy Create free account
hub / github.com/FunAudioLLM/FunMusic / forward

Method forward

inspiremusic/flow/decoder.py:202–277  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 1

upsampleFunction · 0.85

Tested by

no test coverage detected