Conversation
…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>
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
What does this PR do ?
Makes
SortformerEncLabelModel.process_signalscale 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-wideaudio_signal.max()with a per-rowamax(dim=1)over the samples belowaudio_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) inTestSortformerEncLabelModelOffline, next to the existingprocess_signal/ OOM-safe feature extraction tests. It builds its model with_create_sortformer_model.The problem
Line numbers in this description refer to
mainat cf724ac. Onmain, lines 873-875:audio_signalis(batch, num_samples), and.max()has nodim. Every row is therefore divided by the single largest sample in the batch, padding included.audio_signal_lengthis 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:mainOn
mainthe 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'snormalizesetting. 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:normalizemainmainNA(as insortformer_offline_8spk.yaml)per_feature(as insortformer_diarizer_hybrid_loss_4spk-v1.yaml)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. Withper_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) onmain, and 6e-08 with this PR. For the 0.3 file it is 7.1e-06 and 0.021 onmain, 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
InferenceProfilerinstalls around it, the only caller ofprocess_signalinnemo/andexamples/isforward(L909). The scaling above is skipped whenstreaming_modeis true (L874), so it affects every non-streaming path:training_step(L1712, callsforwardat L1729)validation_step(L1859, callsforwardat L1882).test_stepdelegates to it.test_batch(L2005, callsforwardat L2028), used bye2e_diarize_speech.pydiarize(), through_diarize_forward(L678, callsforwardat L692)Inference defaults to
batch_size=1in bothdiarize()ande2e_diarize_speech.py, and an unpadded batch of one is unchanged (see below). The bug appears when inference runs withbatch_size > 1, and in training and validation. The in-repo non-streaming configs train withbatch_size: 8(sortformer_diarizer_hybrid_loss_4spk-v1.yaml, which leavesstreaming_modeat its defaultFalse) andbatch_size: 128(sortformer_offline_8spk.yaml, which also validates and tests withbatch_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 onmainto +0.9804 with this change. Training input is then scaled the waybatch_size=1inference already scales it.Two consequences that you are better placed to weigh than I am:
mainwith a batch size above one has seen the batch-wide scale as a batch-dependent gain on its quieter rows. Undernormalize: NAthat 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.sortformer_offline_8spk.yamlvalidates withbatch_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 inforward_inferand did not touch it.If you would rather keep training bitwise unchanged, applying the per-row reduction only when
not self.trainingis 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 matchbatch_size=1inference, but I am happy to switch.An unpadded batch of one is bitwise unchanged
With no padding, the per-row
amaxover valid samples returns the same value as.max(), and the same operations follow. I checked this withtorch.equalagainst outputs saved frommain(same weights), for the preprocessor input, the features, the feature lengths and the predictions. Four batch-of-one inputs × bothnormalizesettings 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_lengthwould have set the max:mainthen 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 onmain, 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, notmax(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 whatmainalready 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 onmain. The same row on its own is already non-finite onmain.Tests
test_process_signal_peak_normalization_is_per_rowhas 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 pastaudio_signal_length. The existing OOM-safe tests also passrandnbeyond 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:(1 / (x.max() + eps)) * xbitwise, which is today's batch-of-one behaviour.loud_row, the loud row's input equals(1 / (loud.max() + eps)) * loudbitwise.rtol=atol=1e-5).On
mainthe 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. Onmainit 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.I also ran
test_diar_sortformer_modules.py,utils/test_sortformer_utils.py,mixins/test_diarization.pyandutils/test_multispk_instance_manager.pytogether with the file above. The only difference with and without this change is the two new tests:2 failed, 327 passed, 11 skippedversus329 passed, 11 skipped. In both runs, 8 tests intest_multispk_instance_manager.pyerrored because they need thebert-base-casedtokenizer from the Hugging Face Hub, which was unavailable offline.black==24andisortclean at line length 119.Before your PR is "Ready for review"
Pre checks:
torchonlyPR Type:
Additional Information
sortformer_diar_models.pyor its tests (Independent and placeholder-PEE speaker encoder integration for SALMAutomodel + fixes #16088, Add raw-audio streaming sessions for Sortformer #16174, Add per-stream speaker limits to Sortformer sessions #16210, Add CTC-timestamp head to SALM automodel #16259, Add NVFP4 quantized inference for Streaming Sortformer #16276, [core] Unify validation_step_outputs to always return list-of-lists #15470). None of them changesprocess_signal's normalization. Add raw-audio streaming sessions for Sortformer #16174, Add per-stream speaker limits to Sortformer sessions #16210 and Add NVFP4 quantized inference for Streaming Sortformer #16276 also add tests totest_diar_sortformer_models.py, so whichever lands second may need a trivial rebase there.