Skip to content

fix(asr): peak-normalize each Sortformer input row by its own valid samples - #16311

Open
kzos wants to merge 1 commit into
NVIDIA-NeMo:mainfrom
kzos:fix/sortformer-per-row-peak-normalization
Open

kzos wants to merge 1 commit into
NVIDIA-NeMo:mainfrom
kzos:fix/sortformer-per-row-peak-normalization

Conversation

@kzos

@kzos kzos commented Sep 28, 2026

Copy link
Copy Markdown
Contributor

What does this PR do ?

Makes SortformerEncLabelModel.process_signal scale each row of a batch by that row's own peak over its valid samples, instead of by one value taken over the whole padded batch. A file's input to the preprocessor then no longer depends on which other files share its batch.

This also changes what non-streaming training and validation see whenever the batch size is above one. That is spelled out below, together with an eval-only variant if you would rather keep training unchanged.

Collection: ASR

Changelog

  • SortformerEncLabelModel.process_signal: replace the batch-wide audio_signal.max() with a per-row amax(dim=1) over the samples below audio_signal_length. This is on the non-streaming path only. The streaming path does not normalize here and is unchanged. Step 2 of the docstring now says this.
  • tests/collections/speaker_tasks/test_diar_sortformer_models.py: one parametrized test (two cases) in TestSortformerEncLabelModelOffline, next to the existing process_signal / OOM-safe feature extraction tests. It builds its model with _create_sortformer_model.

The problem

Line numbers in this description refer to main at cf724ac. On main, lines 873-875:

        audio_signal_length = audio_signal_length.to(self.device)
        if not self.streaming_mode:
            audio_signal = (1 / (audio_signal.max() + self.eps)) * audio_signal

audio_signal is (batch, num_samples), and .max() has no dim. Every row is therefore divided by the single largest sample in the batch, padding included. audio_signal_length is moved to the device two lines earlier, but the normalization does not use it. With an unpadded batch of one this is per-file peak normalization. In a larger batch, a quiet file next to a loud one gets scaled down by the loud file's peak.

I measured this on CPU with _create_sortformer_model() (random init, eval mode). The quiet file is 3 s of noise whose largest sample is 0.05. The loud file is 4 s of noise whose largest sample is 0.9. The table shows the quiet file's valid samples as handed to the preprocessor:

quiet file min max
alone -1.0359 +0.9804
batched with the loud file, main -0.0586 +0.0555
batched with the loud file, this PR -1.0359 +0.9804

On main the two scales differ by a factor of (0.9 + eps) / (0.05 + eps) = 17.67. How much of that reaches the features depends on the preprocessor's normalize setting. Values are the max |alone − batched| over the quiet file's valid feature frames. The last two columns repeat the run with a milder quiet file whose largest sample is 0.3, a factor of 2.99:

preprocessor normalize 0.05 vs 0.9, main 0.05 vs 0.9, this PR 0.3 vs 0.9, main 0.3 vs 0.9, this PR
NA (as in sortformer_offline_8spk.yaml) 5.743 0 2.193 0
per_feature (as in sortformer_diarizer_hybrid_loss_4spk-v1.yaml) 0.566 7.2e-07 0.0243 7.2e-07

With NA, the difference is 2·ln of the gain ratio (2·ln(17.67) = 5.743, 2·ln(2.99) = 2.193), which is the log-mel shift you get from a gain change. With per_feature, mean subtraction removes most of that shift but not all of it. What remains comes from the additive log zero guard (2^-24), which does not scale with the signal. Feeding the same file at the two gains, one at a time, gives the same 0.566. With the guard set to 1e-30 that drops to 6.2e-05.

With this random-init model the difference also reaches the output probabilities. For the 0.05 file, max |alone − batched| is 2.2e-04 (per_feature) and 0.038 (NA) on main, and 6e-08 with this PR. For the 0.3 file it is 7.1e-06 and 0.021 on main, and at most 1.2e-07 with this PR. A random-init model says nothing about how much DER changes on a trained checkpoint. I have not measured DER on a released checkpoint.

Where this runs

Apart from the timing wrapper that InferenceProfiler installs around it, the only caller of process_signal in nemo/ and examples/ is forward (L909). The scaling above is skipped when streaming_mode is true (L874), so it affects every non-streaming path:

  • training_step (L1712, calls forward at L1729)
  • validation_step (L1859, calls forward at L1882). test_step delegates to it.
  • test_batch (L2005, calls forward at L2028), used by e2e_diarize_speech.py
  • diarize(), through _diarize_forward (L678, calls forward at L692)

Inference defaults to batch_size=1 in both diarize() and e2e_diarize_speech.py, and an unpadded batch of one is unchanged (see below). The bug appears when inference runs with batch_size > 1, and in training and validation. The in-repo non-streaming configs train with batch_size: 8 (sortformer_diarizer_hybrid_loss_4spk-v1.yaml, which leaves streaming_mode at its default False) and batch_size: 128 (sortformer_offline_8spk.yaml, which also validates and tests with batch_size: 32).

This also changes training input for non-streaming models

Training and validation use the same line, so with a batch size above one, each training row is now scaled by its own peak instead of the batch peak. In train() mode, the quiet row's largest sample handed to the preprocessor goes from +0.0555 on main to +0.9804 with this change. Training input is then scaled the way batch_size=1 inference already scales it.

Two consequences that you are better placed to weigh than I am:

  • A model trained on main with a batch size above one has seen the batch-wide scale as a batch-dependent gain on its quieter rows. Under normalize: NA that gain goes straight into the log-mel features (the 5.743 above). This change removes it from training, so the training input distribution changes. I have no DER numbers to show whether that helps or hurts.
  • Validation metrics logged during training move as well, since sortformer_offline_8spk.yaml validates with batch_size: 32.

I could not find the intent of the batch-wide reduction written down anywhere. The expression has been unchanged since the file was added in #11282. #13201 moved it under if not self.streaming_mode, and #16227 only changed the lines around it. #12047 ("fix the issue during batched inference") fixed masking in forward_infer and did not touch it.

If you would rather keep training bitwise unchanged, applying the per-row reduction only when not self.training is a small change and still fixes batched inference. Validation runs in eval mode, so under that variant validation would still switch to per-row scaling. I have not run that variant. I went with the unconditional version because it is simpler and makes training match batch_size=1 inference, but I am happy to switch.

An unpadded batch of one is bitwise unchanged

With no padding, the per-row amax over valid samples returns the same value as .max(), and the same operations follow. I checked this with torch.equal against outputs saved from main (same weights), for the preprocessor input, the features, the feature lengths and the predictions. Four batch-of-one inputs × both normalize settings were all identical: unit-variance noise, a 0.05-peak file, a 0.9-peak file, and a file zero-padded from 2 s to 3 s whose valid part has a positive max.

A batch of one changes only when something past audio_signal_length would have set the max:

  • non-zero padding that is louder than the valid samples, or
  • zero padding after a valid part with no positive sample. main then scales the row by 1/eps (×1000), because the max is a padding zero. A 3 s file with valid samples in [-0.25, -0.2], zero-padded to 4 s, comes out as -249.9994 to -200.0000 on main, and +1.0050 to +1.2563 with this PR.

In both cases the valid part now comes out exactly as it does for the same file without padding (torch.equal).

The reduction is still max, not max(abs), exactly as before. I left that alone, but two consequences are worth stating. A row can land slightly outside [-1, 1] (-1.0359 above). And a row with no positive valid sample is scaled by its own negative max, which is what main already does for that row on its own, but which a positive peak elsewhere in the batch used to mask. If that max is exactly -eps, the scale is infinite. A constant -1e-3 row batched with a loud row gives non-finite features for that row with this change and finite ones on main. The same row on its own is already non-finite on main.

Tests

test_process_signal_peak_normalization_is_per_row has two cases:

  • loud_row: the quiet row is zero-padded and batched with a louder row.
  • loud_padding: a batch of one whose samples past the valid length are louder than its valid part. The case pins that the reduction ignores whatever lies past audio_signal_length. The existing OOM-safe tests also pass randn beyond the valid lengths.

The quiet row is offset so that its most negative sample outweighs its largest one, so a max(abs) reduction would not pass. Each case asserts:

  • Alone, the quiet row's input to the preprocessor equals (1 / (x.max() + eps)) * x bitwise, which is today's batch-of-one behaviour.
  • Inside the batch, it equals the same thing bitwise.
  • In loud_row, the loud row's input equals (1 / (loud.max() + eps)) * loud bitwise.
  • Inside the batch, the quiet row's features match the alone features (rtol=atol=1e-5).

On main the in-batch waveform assertion fails in both cases. I also removed the in-batch waveform assertions temporarily to check that the feature assertion is not redundant. On main it fails on its own, with max absolute differences of 0.1009 and 0.0679.

I also ran the test against some plausible wrong fixes. Each of these fails at least one case: ignoring the length, a batch-wide max over valid samples, max(abs), an off-by-one mask, a flipped eps sign, and using row 0's peak for every row. Filling the padding with 0 instead of -inf is not caught, because the two differ only on a row with no positive valid sample.

tests/collections/speaker_tasks/test_diar_sortformer_models.py
without this change : 2 failed, 201 passed
with this change    : 203 passed

I also ran test_diar_sortformer_modules.py, utils/test_sortformer_utils.py, mixins/test_diarization.py and utils/test_multispk_instance_manager.py together with the file above. The only difference with and without this change is the two new tests: 2 failed, 327 passed, 11 skipped versus 329 passed, 11 skipped. In both runs, 8 tests in test_multispk_instance_manager.py errored because they need the bert-base-cased tokenizer from the Hugging Face Hub, which was unavailable offline.

black==24 and isort clean at line length 119.

Before your PR is "Ready for review"

Pre checks:

  • Make sure you read and followed Contributor guidelines
  • Did you write any new necessary tests?
  • Did you add or update any necessary documentation?
  • Does the PR affect components that are optional to install? (Ex: Numba, Pynini, Apex etc) — no, torch only

PR Type:

  • New Feature
  • Bugfix
  • Documentation

Additional Information

…amples

SortformerEncLabelModel.process_signal scaled the whole (batch, samples)
tensor by audio_signal.max(), one value taken over every row and every
padded sample. The waveform a file hands to the preprocessor, and so its
features and predictions, depended on which other files shared its
batch. Take the max per row over each row's valid samples instead.

A batch of one without padding is scaled exactly as before. Streaming
mode does not normalize here and is unaffected. Offline training and
validation with batch size > 1 go through the same path and now also
see per-row scaling.

Signed-off-by: Zaheer Sheriff K <zaheersheriff.k@gmail.com>
@copy-pr-bot

copy-pr-bot Bot commented Sep 28, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@github-actions github-actions Bot added the ASR label Sep 28, 2026
@svcnvidia-nemo-ci svcnvidia-nemo-ci added the waiting-on-maintainers Waiting on maintainers to respond label Sep 30, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

ASR community-request waiting-on-maintainers Waiting on maintainers to respond

2 participants