multimodalart HF Staff Claude Opus 4.6 (1M context) commited on
Commit
4e71cc6
·
1 Parent(s): a2694ec

Fix audio distortion by normalizing to -1 dBFS with headroom

Browse files

Peak-normalize diffusion output instead of hard-clamping to [-1,1],
which was chopping peaks and causing audible clipping distortion.
Also save as float32 WAV to let torchaudio handle encoding safely.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>

Files changed (1) hide show
  1. app.py +16 -15
app.py CHANGED
@@ -132,18 +132,20 @@ def generate_audio(prompt, negative_prompt, bars, bpm, note, scale, steps, cfg_s
132
 
133
  # 1. Rearrange to [channels, samples] and ensure float32
134
  audio = rearrange(audio, "b d n -> d (b n)").to(torch.float32)
135
-
136
- # 2. PEAK NORMALIZATION (The RC Repo Fix)
137
- # This scales the entire clip so the highest peak is exactly 1.0/-1.0.
138
- # This prevents the 'overdriven' sound by keeping all data within the valid range.
 
 
139
  max_amp = torch.abs(audio).max()
140
- if max_amp > 0:
141
- audio = audio / max_amp
142
-
143
  # 3. Trim to the exact deterministic grid length
144
  end = min(int(audio.shape[-1]), int(clip_samples))
145
  audio = audio[:, :max(1, end)].contiguous()
146
-
147
  # 4. Apply a tiny 15ms fade-out to avoid clicks at the end
148
  fade_ms = 15.0
149
  fade_len = int(round((fade_ms / 1000.0) * SAMPLE_RATE))
@@ -151,15 +153,14 @@ def generate_audio(prompt, negative_prompt, bars, bpm, note, scale, steps, cfg_s
151
  fade_len = min(fade_len, audio.shape[-1])
152
  ramp = torch.linspace(1.0, 0.0, steps=fade_len, device=audio.device)
153
  audio[:, -fade_len:] *= ramp
154
-
155
- # 5. Convert to 16-bit PCM WAV (Standard for Stable Audio outputs)
156
- # We clamp just to be safe, though normalization should already have handled it.
157
- wav_i16 = (audio.clamp(-1, 1) * 32767.0).to(torch.int16).cpu()
158
-
159
  with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
160
  output_path = tmp.name
161
-
162
- torchaudio.save(output_path, wav_i16, SAMPLE_RATE)
163
  return output_path
164
 
165
  except Exception as e:
 
132
 
133
  # 1. Rearrange to [channels, samples] and ensure float32
134
  audio = rearrange(audio, "b d n -> d (b n)").to(torch.float32)
135
+
136
+ # 2. PEAK NORMALIZATION — prevents clipping distortion
137
+ # The diffusion model outputs can far exceed [-1, 1].
138
+ # Hard-clamping (the repo default) chops those peaks → audible distortion.
139
+ # Instead, scale the whole waveform so the loudest peak sits at -1 dBFS
140
+ # (≈0.89), leaving headroom to avoid inter-sample clipping in DAWs.
141
  max_amp = torch.abs(audio).max()
142
+ if max_amp > 1e-8:
143
+ audio = audio / max_amp * 0.89125 # -1 dBFS headroom
144
+
145
  # 3. Trim to the exact deterministic grid length
146
  end = min(int(audio.shape[-1]), int(clip_samples))
147
  audio = audio[:, :max(1, end)].contiguous()
148
+
149
  # 4. Apply a tiny 15ms fade-out to avoid clicks at the end
150
  fade_ms = 15.0
151
  fade_len = int(round((fade_ms / 1000.0) * SAMPLE_RATE))
 
153
  fade_len = min(fade_len, audio.shape[-1])
154
  ramp = torch.linspace(1.0, 0.0, steps=fade_len, device=audio.device)
155
  audio[:, -fade_len:] *= ramp
156
+
157
+ # 5. Save as float32 WAV — let torchaudio handle bit-depth encoding
158
+ audio = audio.clamp(-1, 1).cpu()
159
+
 
160
  with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
161
  output_path = tmp.name
162
+
163
+ torchaudio.save(output_path, audio, SAMPLE_RATE)
164
  return output_path
165
 
166
  except Exception as e: