asteroid-team / asteroid-team/asteroid
get_metrics returns a result not consistent with PITLossWrapper
- Dominant language
- Python
- Stars
- 2.6k
- Forks
- 450
- PR merge metrics
- No merged PRs in 30d
Description
## 🐛 Bug
`get_metrics ` returns a result not consistent with `PITLossWrapper`
### To Reproduce
```
from asteroid.metrics import get_metrics
import torch
from asteroid.losses import PITLossWrapper,pairwise_neg_sisdr
while True:
sources = torch.randn(1, 2, 64000)
est_sources = torch.randn(1, 2, 64000)
mixture = torch.randn(64000)
m = get_metrics(mixture.numpy(), sources.numpy()[0, ], est_sources.numpy()[0, ], compute_permutation=True, metrics_list=['si_sdr'])
print(m, " compute_permutation=True")
m = get_metrics(mixture.numpy(), sources.numpy()[0, ], est_sources.numpy()[0, ], compute_permutation=False, metrics_list=['si_sdr'])
print(m, " compute_permutation=False")
asteroid_pit = PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx')
loss_val = asteroid_pit(est_sources, sources)
print(-loss_val.item(), " PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx')")
print()
```
### My Results
```
{'input_si_sdr': -64.04041290283203, 'si_sdr': -57.542795181274414} compute_permutation=True
{'input_si_sdr': -64.04041290283203, 'si_sdr': -51.0589714050293} compute_permutation=False
-51.0589714050293 PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx')
{'input_si_sdr': -61.0158634185791, 'si_sdr': -61.65677261352539} compute_permutation=True
{'input_si_sdr': -61.0158634185791, 'si_sdr': -45.297006607055664} compute_permutation=False
-45.29700469970703 PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx')
{'input_si_sdr': -47.6236515045166, 'si_sdr': -52.26988410949707} compute_permutation=True
{'input_si_sdr': -47.6236515045166, 'si_sdr': -62.26840019226074} compute_permutation=False
-52.2698860168457 PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx')
{'input_si_sdr': -58.1720085144043, 'si_sdr': -57.48860168457031} compute_permutation=True
{'input_si_sdr': -58.1720085144043, 'si_sdr': -57.48860168457031} compute_permutation=False
-50.58534240722656 PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx')
```
### Expected behavior
The `si_sdr` of `compute_permutation=True ` should be the same with `PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx')`?
The `si_sdr` of `compute_permutation=True ` should always be higher than `compute_permutation=False`?
#### Additional info
I don't know if it's my misunderstanding or not.
Contributor guide
Research direction
Start at asteroid.metrics.get_metrics and compare its compute_permutation behavior with PITLossWrapper(pairwise_neg_sisdr, pit_from='pw_mtx') using the reproduction in the issue. Determine the intended permutation semantics, then make the behavior consistent and add regression coverage for both permutation settings.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- machine-learning
- Issue type
- Bug
- Difficulty
- 3/5
- Estimated time
- 1-2 days
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 35/100