asteroid-team / asteroid-team/asteroid

get_metrics returns a result not consistent with PITLossWrapper

Open
#509 14 comments 0 reactions 0 assignees View on GitHub
bug help wanted
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

Open the contributing 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.