Repository navigation
Fix PESQ aborting a whole batch on one unscorable sample - #3455
Conversation
The pesq backend defaults to on_error=PesqError.RAISE_EXCEPTION, so a single degenerate sample (e.g. a silent reference, or a prediction whose amplitude dwarfs the reference after the backend's shared normalisation) raised NoUtterancesError out of the whole update. Pass on_error=PesqError.RETURN_VALUES at every call site and map the returned error codes, as well as the exception objects pesq_batch collects from its workers, onto nan. The result now keeps one entry per input sample, so the documented shape contract holds again for inputs with more than one batch dimension. The class metric leaves nan samples out of its average instead of propagating them. Fixes Lightning-AI#3304
e87534b to
8426d03
Compare
There was a problem hiding this comment.
🟢 Approval recommended
The changes align with the PR description, add targeted regression tests for the reported failure modes, and the only feedback is a minor docstring wording nit.
Pull request overview
This PR updates TorchMetrics’ PESQ functional and class-based metric behavior so that an unscorable sample (e.g., silent reference triggering NoUtterancesError) does not abort or shrink a batch, instead producing nan for that sample while keeping output shape consistent and excluding nan samples from the class metric average.
Changes:
- Pass
on_error=PesqError.RETURN_VALUESto PESQ backend calls and map backend failures tonanwhile preserving one output per input sample. - Restore multidimensional output shaping to
preds.shape[:-1]for the functional interface. - Add regression tests covering degenerate samples (single + batched, single + multi-process) and multidimensional shape behavior; update changelog entry.
File summaries
| File | Description |
|---|---|
| tests/unittests/audio/test_pesq.py | Adds tests ensuring degenerate samples yield nan without aborting batches and that multidim inputs preserve shape. |
| src/torchmetrics/functional/audio/pesq.py | Ensures backend errors become nan (not exceptions/dropped entries) and restores documented output shaping. |
| src/torchmetrics/audio/pesq.py | Updates class metric aggregation to ignore nan samples in sum/total. |
| CHANGELOG.md | Documents the PESQ degenerate-sample behavior fix under Unreleased. |
Review details
- Files reviewed: 4/4 changed files
- Comments generated: 1
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
|
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## master #3455 +/- ##
========================================
- Coverage 37% 31% -6%
========================================
Files 349 349
Lines 19901 19899 -2
========================================
- Hits 7264 6157 -1107
- Misses 12637 13742 +1105 🚀 New features to boost your workflow:
|
There was a problem hiding this comment.
🟢 Approval recommended
The change is narrowly scoped, aligns behavior across single/multiprocess paths, preserves documented shapes, and is backed by targeted regression tests for the previously failing scenarios.
Review details
- Files reviewed: 4/4 changed files
- Comments generated: 0 new
- Review effort level: Lite
Addresses review feedback on Lightning-AI#3455. All three backend calls were made multi-line by this PR, so name the arguments: the upstream signatures are pesq(fs, ref, deg, mode, on_error) and pesq_batch(fs, ref, deg, mode, n_processor, on_error), which makes the target -> ref and preds -> deg mapping explicit at the call site. Also hyphenates "class-based metric" in the note added by this PR. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_014puKzWNVkoZhuj4LeGZwjD
|
Tick the box to add this pull request to the merge queue (same as
|
|
Fixes #3304.
One unscorable sample takes down the whole batch.
pesqandpesq_batchbothdefault to
on_error=PesqError.RAISE_EXCEPTIONand torchmetrics never passeson_error, so what happens depends onn_processes:n_processes=1(the default):pesq()raisesNoUtterancesErrorand itpropagates out of
update(). The_filter_error_msgguard added in Ignore theNoUtterancesErrorwhen calculating pesq for a batch #2753 isunreachable on this path.
n_processes != 1:pesq_batch()catches worker exceptions and returns themin the result list, so the guard is reached, but it drops the failed
entry. A 3-sample batch silently returns 2 values.
Measured on a 3-sample batch with
target[1]silent:After:
Change
Pass
on_error=PesqError.RETURN_VALUESat all three call sites and map everyfailure report onto
nan, keeping one entry per input sample. Failures arrivetwo ways and both are handled: negative error codes (
-1..-7, which cannotcollide with a real score since valid PESQ is
>= -0.5) and the exceptionobjects
pesq_batchcollects from its workers.The class metric excludes
nansamples from both the running sum andtotal,so a degenerate sample does not poison an epoch's average. That matches what the
multiprocessing path already did by dropping them, so it is not a new policy.
Scores for scorable samples are unchanged:
RETURN_VALUESandRAISE_EXCEPTIONreturn the same value on success (verified,2.158228635787964 both ways).
Also in this PR
#2753 replaced
pesq_val.reshape(preds.shape[:-1])withreshape(len(pesq_val)), so a(2, 3, 2100)input returned(6,)instead of(2, 3), contradicting the documented(...,)shape. Restored, which is safenow that no entries are dropped. This is a behaviour change for
ndim > 2inputs relative to 1.4.x-1.9.0. Happy to split it into its own PR if you would
rather keep this one minimal.
Tests
Four cases in
tests/unittests/audio/test_pesq.py, all failing on main:The two different failure modes for
[1]and[2]are the two paths above.tests/unittests/audio/test_pesq.py: 20 -> 24 passed, same 6 skipped and 3xfailed. Doctests pass. Wider audio dir: 88 passed. ruff, format and mypy clean.