Views
No views yet
from dynamic_source_separator import DynamicSourceSeparator
model = DynamicSourceSeparator.from_pretrained(
"cadenzachallenge/Dynamic_Source_Separator_Causal"
).cpu()
1def dynamic_masked_loss(mixture, separated_sources, ground_truth_sources, indicator):
2 # Reconstruction Loss
3 reconstruction = sum(separated_sources.values())
4 reconstruction_loss = nn.L1Loss()(reconstruction, mixture)
5 # Separation Loss
6 separation_loss = 0
7 for instrument, active in indicator.items():
8 if active:
9 separation_loss += nn.L1Loss()(
10 separated_sources[instrument], ground_truth_sources[instrument]
11 )
12 return reconstruction_loss + separation_loss