diff --git a/README.md b/README.md index f33b1cd..b41c2a5 100644 --- a/README.md +++ b/README.md @@ -15,12 +15,14 @@ workflows, especially when dealing with multiple audio inputs or outputs. - [6. Audio Channel Conv and Resampler](#6-audio-channel-conv-and-resampler) - [7. Audio Information](#7-audio-information) - [8. Audio Cut](#8-audio-cut) - - [9. Audio Concatenate](#8-audio-concatenate) - - [10. Audio Blend](#9-audio-blend) - - [11. Audio Test Signal Generator](#10-audio-test-signal-generator) - - [12. Audio Musical Note](#11-audio-musical-note) - - [13. Audio Join 2 Channels](#12-audio-join-2-channels) - - [14. Audio Split 2 Channels](#13-audio-split-2-channels) + - [9. Audio Concatenate](#9-audio-concatenate) + - [10. Audio Blend](#10-audio-blend) + - [11. Audio Test Signal Generator](#11-audio-test-signal-generator) + - [12. Audio Musical Note](#12-audio-musical-note) + - [13. Audio Join 2 Channels](#13-audio-join-2-channels) + - [14. Audio Split 2 Channels](#14-audio-split-2-channels) + - [15. Audio Normalize (Peak)](#15-audio-normalize-peak) + - [16. Audio Apply Batched Gain](#16-audio-apply-batched-gain) - [🚀 Installation](#-installation) - [📦 Dependencies](#-dependencies) - [🖼️ Examples](#️-examples) @@ -259,6 +261,30 @@ workflows, especially when dealing with multiple audio inputs or outputs. - `audio_left` (AUDIO): A mono audio signal containing only the left channel data. - `audio_right` (AUDIO): A mono audio signal containing only the right channel data. +### 15. Audio Normalize (Peak) + - **Display Name:** `Audio Normalize (Peak)` + - **Internal Name:** `SET_AudioNormalize` + - **Category:** `audio/manipulation` + - **Description:** Normalizes the volume of an audio signal so that its loudest point (peak) reaches a specified target level. This is useful for maximizing volume without clipping. + - **Inputs:** + - `audio` (AUDIO): The audio to normalize. Supports batches. + - `peak_level` (FLOAT): The target peak amplitude level. `1.0` is the maximum possible level (0 dBFS). Normalizing to slightly less, like `0.9`, can provide headroom. + - **Outputs:** + - `normalized_audio` (AUDIO): The audio with its volume adjusted. + - `original_peak_level` (FLOAT): The original peak level of the input audio for each item in the batch. **This value can be used with the `Audio Apply Batched Gain` node to revert the normalization and restore the original volume.** + +### 16. Audio Apply Batched Gain + - **Display Name:** `Audio Apply Batched Gain` + - **Internal Name:** `SET_AudioApplyBatchedGain` + - **Category:** `audio/manipulation` + - **Description:** Applies a separate gain (volume) level to each item in an audio batch. This is the perfect companion to the `Audio Normalize` node for reverting normalization. + - **Inputs:** + - `audio` (AUDIO): The audio batch to apply gain to. + - `gain_values` (FLOAT): A batch of gain values. The node expects this to be a 1D tensor of shape `(batch_size,)`, which is the format provided by the `original_peak_level` output of the `Audio Normalize` node. + - **Output:** + - `audio_out` (AUDIO): The audio with the per-item gain applied. + + ## 🚀 Installation You can install the nodes from the ComfyUI nodes manager, the name is *Audio Batch*, or just do it manually: @@ -291,6 +317,7 @@ Once installed the examples are available in the ComfyUI workflow templates, in - [generate_and_blend.json](example_workflows/generate_and_blend.json): Shows how to generate four musical notes and blend them together to create a chord. - [cut_and_concat.json](example_workflows/cut_and_concat.json): Shows how to cut and concatenate audio. +- [normalize_and_undo.json](example_workflows/normalize_and_undo.json): Shows how to normalize audio level and then revert it. ## 📝 Usage Notes diff --git a/example_workflows/normalize_and_undo.jpg b/example_workflows/normalize_and_undo.jpg new file mode 100644 index 0000000..d93e135 Binary files /dev/null and b/example_workflows/normalize_and_undo.jpg differ diff --git a/example_workflows/normalize_and_undo.json b/example_workflows/normalize_and_undo.json new file mode 100644 index 0000000..13debab --- /dev/null +++ b/example_workflows/normalize_and_undo.json @@ -0,0 +1 @@ +{"id":"e19f3025-8818-483c-a738-6156d8f94f9a","revision":0,"last_node_id":15,"last_link_id":11,"nodes":[{"id":1,"type":"LoadAudio","pos":[1702.73046875,1176.981689453125],"size":[300,140],"flags":{},"order":0,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"COMBO","widget":{"name":"audio"},"link":null},{"localized_name":"audioUI","name":"audioUI","type":"AUDIO_UI","widget":{"name":"audioUI"},"link":null},{"localized_name":"choose file to upload","name":"upload","type":"AUDIOUPLOAD","widget":{"name":"upload"},"link":null}],"outputs":[{"localized_name":"AUDIO","name":"AUDIO","type":"AUDIO","links":[2,6]}],"properties":{"cnr_id":"comfy-core","ver":"0.3.34","Node name for S&R":"LoadAudio","widget_ue_connectable":{}},"widgets_values":["aud-sample-vs-1.wav",null,null],"color":"#222","bgcolor":"#000"},{"id":3,"type":"SET_AudioNormalize","pos":[2059.697998046875,1175.7960205078125],"size":[264.7818908691406,78],"flags":{},"order":5,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"AUDIO","link":2},{"localized_name":"peak_level","name":"peak_level","type":"FLOAT","widget":{"name":"peak_level"},"link":null}],"outputs":[{"localized_name":"normalized_audio","name":"normalized_audio","type":"AUDIO","links":[1,4,11]},{"localized_name":"original_peak_level","name":"original_peak_level","type":"TORCH_TENSOR","links":[3]}],"properties":{"aux_id":"set-soft/ComfyUI-AudioBatch","ver":"3286ba248c48661cb0461d45b3fca5ff0f40f26a","Node name for S&R":"SET_AudioNormalize"},"widgets_values":[1],"color":"#233","bgcolor":"#355"},{"id":4,"type":"SET_AudioApplyBatchedGain","pos":[2472.402587890625,1352.4998779296875],"size":[203.34042358398438,46],"flags":{},"order":9,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"AUDIO","link":4},{"localized_name":"gain_values","name":"gain_values","type":"TORCH_TENSOR","link":3}],"outputs":[{"localized_name":"audio_out","name":"audio_out","type":"AUDIO","links":[5,9]}],"properties":{"aux_id":"set-soft/ComfyUI-AudioBatch","ver":"3286ba248c48661cb0461d45b3fca5ff0f40f26a","Node name for S&R":"SET_AudioApplyBatchedGain"},"widgets_values":[],"color":"#233","bgcolor":"#355"},{"id":2,"type":"PreviewAudio","pos":[2419.034912109375,1169.8660888671875],"size":[344.231201171875,95.92967224121094],"flags":{},"order":7,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"AUDIO","link":1},{"localized_name":"audioUI","name":"audioUI","type":"AUDIO_UI","widget":{"name":"audioUI"},"link":null}],"outputs":[],"properties":{"cnr_id":"comfy-core","ver":"0.3.34","Node name for S&R":"PreviewAudio","widget_ue_connectable":{}},"widgets_values":[],"color":"#222","bgcolor":"#000"},{"id":11,"type":"easy showAnything","pos":[2608.785400390625,1012.1365356445312],"size":[210,88],"flags":{},"order":11,"mode":0,"inputs":[{"localized_name":"anything","name":"anything","shape":7,"type":"*","link":10}],"outputs":[{"localized_name":"output","name":"output","type":"*","links":null}],"properties":{"cnr_id":"comfyui-easy-use","ver":"1.3.1","Node name for S&R":"easy showAnything"},"widgets_values":["tensor(1.)"],"color":"#222","bgcolor":"#000"},{"id":10,"type":"SET_AudioInfo","pos":[2433.266845703125,957.5831298828125],"size":[149.63046264648438,166],"flags":{"collapsed":false},"order":8,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"AUDIO","link":11}],"outputs":[{"localized_name":"audio_bypass","name":"audio_bypass","type":"AUDIO","links":null},{"localized_name":"batch_size","name":"batch_size","type":"INT","links":null},{"localized_name":"channels","name":"channels","type":"INT","links":null},{"localized_name":"num_samples","name":"num_samples","type":"INT","links":null},{"localized_name":"sample_rate","name":"sample_rate","type":"INT","links":null},{"localized_name":"mean","name":"mean","type":"TORCH_TENSOR","links":null},{"localized_name":"std","name":"std","type":"TORCH_TENSOR","links":null},{"localized_name":"peak","name":"peak","type":"TORCH_TENSOR","links":[10]}],"properties":{"aux_id":"set-soft/ComfyUI-AudioBatch","ver":"3286ba248c48661cb0461d45b3fca5ff0f40f26a","Node name for S&R":"SET_AudioInfo"},"widgets_values":[],"color":"#2a363b","bgcolor":"#3f5159"},{"id":8,"type":"SET_AudioInfo","pos":[2787.859130859375,1504.3001708984375],"size":[149.63046264648438,166],"flags":{},"order":13,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"AUDIO","link":9}],"outputs":[{"localized_name":"audio_bypass","name":"audio_bypass","type":"AUDIO","links":null},{"localized_name":"batch_size","name":"batch_size","type":"INT","links":null},{"localized_name":"channels","name":"channels","type":"INT","links":null},{"localized_name":"num_samples","name":"num_samples","type":"INT","links":null},{"localized_name":"sample_rate","name":"sample_rate","type":"INT","links":null},{"localized_name":"mean","name":"mean","type":"TORCH_TENSOR","links":null},{"localized_name":"std","name":"std","type":"TORCH_TENSOR","links":null},{"localized_name":"peak","name":"peak","type":"TORCH_TENSOR","links":[8]}],"properties":{"aux_id":"set-soft/ComfyUI-AudioBatch","ver":"3286ba248c48661cb0461d45b3fca5ff0f40f26a","Node name for S&R":"SET_AudioInfo"},"widgets_values":[],"color":"#2a363b","bgcolor":"#3f5159"},{"id":9,"type":"easy showAnything","pos":[2965.74951171875,1503.114013671875],"size":[210,88],"flags":{},"order":14,"mode":0,"inputs":[{"localized_name":"anything","name":"anything","shape":7,"type":"*","link":8}],"outputs":[{"localized_name":"output","name":"output","type":"*","links":null}],"properties":{"cnr_id":"comfyui-easy-use","ver":"1.3.1","Node name for S&R":"easy showAnything"},"widgets_values":["tensor(0.8055)"],"color":"#222","bgcolor":"#000"},{"id":5,"type":"PreviewAudio","pos":[2767.698974609375,1353.6854248046875],"size":[344.231201171875,95.92967224121094],"flags":{},"order":12,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"AUDIO","link":5},{"localized_name":"audioUI","name":"audioUI","type":"AUDIO_UI","widget":{"name":"audioUI"},"link":null}],"outputs":[],"properties":{"cnr_id":"comfy-core","ver":"0.3.34","Node name for S&R":"PreviewAudio","widget_ue_connectable":{}},"widgets_values":[],"color":"#222","bgcolor":"#000"},{"id":6,"type":"SET_AudioInfo","pos":[1720.5186767578125,961.1417846679688],"size":[149.63046264648438,166],"flags":{},"order":6,"mode":0,"inputs":[{"localized_name":"audio","name":"audio","type":"AUDIO","link":6}],"outputs":[{"localized_name":"audio_bypass","name":"audio_bypass","type":"AUDIO","links":null},{"localized_name":"batch_size","name":"batch_size","type":"INT","links":null},{"localized_name":"channels","name":"channels","type":"INT","links":null},{"localized_name":"num_samples","name":"num_samples","type":"INT","links":null},{"localized_name":"sample_rate","name":"sample_rate","type":"INT","links":null},{"localized_name":"mean","name":"mean","type":"TORCH_TENSOR","links":null},{"localized_name":"std","name":"std","type":"TORCH_TENSOR","links":null},{"localized_name":"peak","name":"peak","type":"TORCH_TENSOR","links":[7]}],"properties":{"aux_id":"set-soft/ComfyUI-AudioBatch","ver":"3286ba248c48661cb0461d45b3fca5ff0f40f26a","Node name for S&R":"SET_AudioInfo"},"widgets_values":[],"color":"#2a363b","bgcolor":"#3f5159"},{"id":7,"type":"easy showAnything","pos":[1888.9237060546875,984.860595703125],"size":[210,88],"flags":{},"order":10,"mode":0,"inputs":[{"localized_name":"anything","name":"anything","shape":7,"type":"*","link":7}],"outputs":[{"localized_name":"output","name":"output","type":"*","links":null}],"properties":{"cnr_id":"comfyui-easy-use","ver":"1.3.1","Node name for S&R":"easy showAnything"},"widgets_values":["tensor(0.8055)"],"color":"#222","bgcolor":"#000"},{"id":12,"type":"MarkdownNote","pos":[1705.695556640625,1372.221923828125],"size":[287.8678894042969,88],"flags":{},"order":1,"mode":0,"inputs":[],"outputs":[],"properties":{},"widgets_values":["# 1. Select an audio \nYou can see the peak value"],"color":"#432","bgcolor":"#653"},{"id":13,"type":"MarkdownNote","pos":[2478.55224609375,802.00537109375],"size":[287.8678894042969,88],"flags":{},"order":2,"mode":0,"inputs":[],"outputs":[],"properties":{},"widgets_values":["# 2. Here you have the audio with peak == 1"],"color":"#432","bgcolor":"#653"},{"id":14,"type":"MarkdownNote","pos":[2817.0810546875,1218.1605224609375],"size":[287.8678894042969,88],"flags":{},"order":3,"mode":0,"inputs":[],"outputs":[],"properties":{},"widgets_values":["# 3. And here the audio is back to its original level"],"color":"#432","bgcolor":"#653"},{"id":15,"type":"MarkdownNote","pos":[2184.859375,1545.54248046875],"size":[428.8094787597656,88],"flags":{},"order":4,"mode":0,"inputs":[],"outputs":[],"properties":{},"widgets_values":["# Show Any is part of \"Easy-Use\" nodes"],"color":"#432","bgcolor":"#653"}],"links":[[1,3,0,2,0,"AUDIO"],[2,1,0,3,0,"AUDIO"],[3,3,1,4,1,"TORCH_TENSOR"],[4,3,0,4,0,"AUDIO"],[5,4,0,5,0,"AUDIO"],[6,1,0,6,0,"AUDIO"],[7,6,7,7,0,"*"],[8,8,7,9,0,"*"],[9,4,0,8,0,"AUDIO"],[10,10,7,11,0,"*"],[11,3,0,10,0,"AUDIO"]],"groups":[],"config":{},"extra":{"ds":{"scale":1.2345539040268603,"offset":[-1579.0805695433937,-1056.3910821825534]}},"version":0.4} diff --git a/source/nodes/nodes_audio.py b/source/nodes/nodes_audio.py index 46cb238..0dbba63 100644 --- a/source/nodes/nodes_audio.py +++ b/source/nodes/nodes_audio.py @@ -904,3 +904,131 @@ class AudioConcatenate: "sample_rate": target_sr } return (output_audio,) + + +class AudioNormalize: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "audio": ("AUDIO", {"tooltip": "The audio to normalize."}), + "peak_level": ("FLOAT", { + "default": 1.0, + "min": 0.0, + "max": 10.0, # Allow some headroom if needed + "step": 0.01, + "tooltip": "The target peak amplitude level. 1.0 is 0 dBFS (maximum)." + }) + }, + } + + RETURN_TYPES = ("AUDIO", "TORCH_TENSOR") + RETURN_NAMES = ("normalized_audio", "original_peak_level") + FUNCTION = "normalize_audio" + CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY + DESCRIPTION = "Normalizes audio so its loudest peak reaches a target level." + UNIQUE_NAME = "SET_AudioNormalize" + DISPLAY_NAME = "Audio Normalize (Peak)" + + def normalize_audio(self, audio: dict, peak_level: float): + waveform = audio['waveform'] + sample_rate = audio['sample_rate'] + + logger.debug(f"Normalizing audio of shape {waveform.shape} to target peak: {peak_level}") + + # Find the maximum absolute value for each batch item. + # This is the original peak level we need to return. + original_peak_level, _ = torch.max(torch.abs(waveform), dim=2, keepdim=True) + original_peak_level, _ = torch.max(original_peak_level, dim=1, keepdim=True) + # original_peak_level shape is (B, 1, 1) + + # Define a threshold for what we consider silent to avoid division by zero + silence_threshold = 1e-8 + + # Calculate the gain factor needed to reach the target peak_level. + # Add the threshold to the denominator to prevent division by zero for silent clips. + gain_factor = peak_level / (original_peak_level + silence_threshold) + + # For silent clips (where original_peak_level is below threshold), the gain + # will be huge but will be multiplied by ~0, resulting in ~0 (silence). + # We can explicitly set gain to 0 for silent clips to be safer. + is_silent = original_peak_level <= silence_threshold + gain_factor[is_silent] = 0.0 + + # Apply the gain + normalized_waveform = waveform * gain_factor + + # The value needed to revert is the original peak level. + # Squeeze it to be a more convenient shape for other nodes (e.g., a 1D tensor or a scalar if B=1) + # Squeeze removes all dims of size 1. (B,1,1) -> (B) + revert_value = original_peak_level.squeeze() + # If the original batch size was 1, this will be a 0-dim tensor (scalar). + # Let's ensure it's at least a 1D tensor for consistency. + if revert_value.ndim == 0: + revert_value = revert_value.unsqueeze(0) # Make it (1,) + + logger.info(f"Peak normalization applied. Original avg peak: {original_peak_level.mean():.4f}, " + f"Target peak: {peak_level}") + + output_audio = {"waveform": normalized_waveform, "sample_rate": sample_rate} + + # The revert_value is the original peak level for each item in the batch + return (output_audio, revert_value) + + +class AudioApplyBatchedGain: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "audio": ("AUDIO", {"tooltip": "The audio batch to apply gain to."}), + "gain_values": ("TORCH_TENSOR", { + "tooltip": "A batch of gain values (e.g., from an Audio Normalize node). Expects a 1D tensor." + }), + }, + } + + RETURN_TYPES = ("AUDIO",) + RETURN_NAMES = ("audio_out",) + FUNCTION = "apply_gain" + CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY + DESCRIPTION = "Applies a different gain (volume) to each item in an audio batch." + UNIQUE_NAME = "SET_AudioApplyBatchedGain" + DISPLAY_NAME = "Audio Apply Batched Gain" + + def apply_gain(self, audio: dict, gain_values: torch.Tensor): + waveform = audio['waveform'] + sample_rate = audio['sample_rate'] + + batch_size = waveform.shape[0] + + # Ensure gain_values is a tensor + if not isinstance(gain_values, torch.Tensor): + # If a simple float is passed, convert it to a tensor + # This makes the node more flexible + gain_values = torch.tensor([gain_values], dtype=torch.float32) + + # Validate dimensions + if gain_values.ndim == 0: # Handle scalar tensor + gain_values = gain_values.unsqueeze(0) + + if gain_values.shape[0] != batch_size: + msg = (f"Batch size of audio ({batch_size}) and gain_values ({gain_values.shape[0]}) must match. " + "Connect the `original_peak_level` from an Audio Normalize node.") + logger.error(msg) + raise ValueError(msg) + + # Reshape gain_values for broadcasting: (B) -> (B, 1, 1) + # This allows multiplying (B, C, N) with (B, 1, 1) + gain_tensor = gain_values.to(waveform.device, waveform.dtype).reshape(batch_size, 1, 1) + + # Apply the per-item gain + adjusted_waveform = waveform * gain_tensor + + logger.info(f"Applied batched gain. Gain range: [{gain_values.min():.4f}, {gain_values.max():.4f}]") + + output_audio = { + "waveform": adjusted_waveform, + "sample_rate": sample_rate + } + return (output_audio,) diff --git a/source/tests/test_audio_normalize.py b/source/tests/test_audio_normalize.py new file mode 100644 index 0000000..a78bacb --- /dev/null +++ b/source/tests/test_audio_normalize.py @@ -0,0 +1,123 @@ +""" +Regression tests for the AudioNormalize and AudioApplyBatchedGain nodes in ComfyUI-AudioBatch. +""" + +import torch +import pytest +import bootstrap # noqa: F401 +from nodes.nodes_audio import AudioNormalize, AudioApplyBatchedGain + + +# Helper function +def create_dummy_audio(batch_size, channels, samples, sr, peak_value=0.5, device='cpu'): + # Create audio that is not normalized to 1.0 to make testing meaningful + waveform = torch.randn(batch_size, channels, samples, device=device) + # Normalize to a known peak value + current_peak, _ = torch.max(torch.abs(waveform), dim=2, keepdim=True) + current_peak, _ = torch.max(current_peak, dim=1, keepdim=True) + waveform = waveform / (current_peak + 1e-9) * peak_value + return {"waveform": waveform, "sample_rate": sr} + + +@pytest.fixture +def normalize_node(): + return AudioNormalize() + + +@pytest.fixture +def apply_gain_node(): + return AudioApplyBatchedGain() + + +# --- Tests for AudioNormalize --- + +def test_normalize_simple(normalize_node): + """Tests if a simple audio clip is normalized to the target peak level.""" + sr = 44100 + original_peak = 0.25 + target_peak = 1.0 + audio = create_dummy_audio(1, 1, sr, sr, peak_value=original_peak) + + (normalized_audio, original_peak_level) = normalize_node.normalize_audio(audio, target_peak) + + # Assertions + new_peak = torch.max(torch.abs(normalized_audio['waveform'])) + assert new_peak.item() == pytest.approx(target_peak, abs=1e-6) + assert original_peak_level.item() == pytest.approx(original_peak, abs=1e-6) + + +def test_normalize_with_batch(normalize_node): + """Tests that each item in a batch is normalized independently.""" + sr = 44100 + # Create a batch with different peak levels + audio1 = create_dummy_audio(1, 2, sr, sr, peak_value=0.5) + audio2 = create_dummy_audio(1, 2, sr, sr, peak_value=0.1) + + batched_waveform = torch.cat((audio1['waveform'], audio2['waveform']), dim=0) + batched_audio = {"waveform": batched_waveform, "sample_rate": sr} + + target_peak = 0.9 + (normalized_audio, original_peak_levels) = normalize_node.normalize_audio(batched_audio, target_peak) + + # Check item 1 + peak1 = torch.max(torch.abs(normalized_audio['waveform'][0])) + assert peak1.item() == pytest.approx(target_peak, abs=1e-6) + assert original_peak_levels[0].item() == pytest.approx(0.5, abs=1e-6) + + # Check item 2 + peak2 = torch.max(torch.abs(normalized_audio['waveform'][1])) + assert peak2.item() == pytest.approx(target_peak, abs=1e-6) + assert original_peak_levels[1].item() == pytest.approx(0.1, abs=1e-6) + + assert original_peak_levels.shape == (2,) + + +def test_normalize_silent_audio(normalize_node): + """Tests that silent audio remains silent and doesn't cause errors.""" + sr = 44100 + silent_audio = {"waveform": torch.zeros(1, 1, sr), "sample_rate": sr} + + (normalized_audio, original_peak_level) = normalize_node.normalize_audio(silent_audio, 1.0) + + # Assert that the output is still silent + assert torch.all(normalized_audio['waveform'] == 0) + # The original peak level should be 0 + assert original_peak_level.item() == 0.0 + + +# --- Tests for AudioApplyBatchedGain --- + +def test_apply_gain_and_revert_normalization(normalize_node, apply_gain_node): + """ + An integration test to verify that applying the `original_peak_level` + from the normalize node correctly reverts the audio to its original state. + """ + sr = 44100 + audio_orig = create_dummy_audio(3, 2, sr, sr, peak_value=0.3) + + # 1. Normalize the audio + (normalized_audio, original_peak_levels) = normalize_node.normalize_audio(audio_orig, 1.0) + + # 2. Apply the original peak levels as gain to the normalized audio + (reverted_audio,) = apply_gain_node.apply_gain(normalized_audio, original_peak_levels) + + # Assertions + # The reverted audio should be almost identical to the original audio + assert reverted_audio['sample_rate'] == audio_orig['sample_rate'] + assert reverted_audio['waveform'].shape == audio_orig['waveform'].shape + # Check that the content is the same (with a small tolerance for float math) + assert torch.allclose(reverted_audio['waveform'], audio_orig['waveform'], atol=1e-6) + + +def test_apply_gain_mismatched_batch_size_error(apply_gain_node): + """ + Tests that an error is raised if the audio batch size and gain values + batch size do not match. + """ + sr = 44100 + audio = create_dummy_audio(3, 1, sr, sr) + # Gain values for only 2 items, but audio has 3 + gain_values = torch.tensor([0.5, 0.5]) + + with pytest.raises(ValueError, match="Batch size of audio .* and gain_values .* must match"): + apply_gain_node.apply_gain(audio, gain_values)