DSPark 1.8.0
Header-only C++20 DSP for real-time and offline audio
Loading...
Searching...
No Matches
SpectralDenoiser.h
1// DSPark - Professional Audio DSP Framework
2// Copyright (c) 2026 Cristian Moresi - MIT License
3
4#pragma once
5
39#include "../Core/AudioBuffer.h"
40#include "../Core/AudioSpec.h"
41#include "../Core/DspMath.h"
42#include "../Core/SpectralProcessor.h"
43#include "../Core/StateBlob.h"
44
45#include <algorithm>
46#include <atomic>
47#include <cmath>
48#include <cstddef>
49#include <cstdint>
50#include <vector>
51
52namespace dspark {
53
60template <FloatType T>
62{
63public:
64 // -- Lifecycle ---------------------------------------------------------------
65
77 void prepare(const AudioSpec& spec, int fftSize = 2048)
78 {
79 if (!spec.isValid()) return;
80 prepared_.store(false, std::memory_order_relaxed);
81 numChannels_ = spec.numChannels;
82 stft_.prepare(spec, fftSize, fftSize / 4);
83 numBins_ = stft_.getNumBins();
84
85 profile_.assign(static_cast<size_t>(numBins_), 0.0f);
86 learnFrames_ = 0;
87 cleanPower_.assign(static_cast<size_t>(numChannels_),
88 std::vector<float>(static_cast<size_t>(numBins_), 0.0f));
89
90 prepared_.store(true, std::memory_order_relaxed);
91 reset();
92 }
93
95 void reset() noexcept
96 {
97 if (!prepared_.load(std::memory_order_relaxed)) return;
98 stft_.reset();
99 for (auto& c : cleanPower_)
100 std::fill(c.begin(), c.end(), 0.0f);
101 callCounter_ = 0;
102 }
103
105 void clearProfile() noexcept
106 {
107 std::fill(profile_.begin(), profile_.end(), 0.0f);
108 learnFrames_ = 0;
109 }
110
111 // -- Parameters (thread-safe) ---------------------------------------------------
112
114 void setLearning(bool learning) noexcept
115 {
116 learning_.store(learning, std::memory_order_relaxed);
117 }
118
121 void setReduction(T db) noexcept
122 {
123 if (!std::isfinite(db)) return;
124 reduction_.store(std::clamp(db, T(0), T(40)), std::memory_order_relaxed);
125 }
126
131 void setThreshold(T factor) noexcept
132 {
133 if (!std::isfinite(factor)) return;
134 threshold_.store(std::clamp(factor, T(1), T(8)), std::memory_order_relaxed);
135 }
136
137 [[nodiscard]] bool getLearning() const noexcept { return learning_.load(std::memory_order_relaxed); }
138 [[nodiscard]] T getReduction() const noexcept { return reduction_.load(std::memory_order_relaxed); }
139 [[nodiscard]] T getThreshold() const noexcept { return threshold_.load(std::memory_order_relaxed); }
140
142 [[nodiscard]] int getLatency() const noexcept { return stft_.getLatency(); }
143
146 [[nodiscard]] std::vector<uint8_t> getState() const
147 {
148 StateWriter w(stateId("DNSE"), 1);
149 // Explicit float casts: the blob stores float, and with T = double the
150 // unqualified write(key, double) would be ambiguous (float/int32/bool).
151 w.write("reduction", static_cast<float>(reduction_.load(std::memory_order_relaxed)));
152 w.write("threshold", static_cast<float>(threshold_.load(std::memory_order_relaxed)));
153 return w.blob();
154 }
155
157 bool setState(const uint8_t* data, size_t size)
158 {
159 StateReader r(data, size);
160 if (!r.isValid() || r.processorId() != stateId("DNSE")) return false;
161 setReduction(static_cast<T>(r.read("reduction", 18.0f)));
162 setThreshold(static_cast<T>(r.read("threshold", 2.0f)));
163 return true;
164 }
165
166 // -- Processing -------------------------------------------------------------------
167
169 void processBlock(AudioBufferView<T> buffer) noexcept
170 {
171 if (!prepared_.load(std::memory_order_relaxed)) return;
172
173 const bool learning = learning_.load(std::memory_order_relaxed);
174 const float floorGain = std::pow(
175 10.0f, -static_cast<float>(reduction_.load(std::memory_order_relaxed)) / 20.0f);
176 const float thresh = static_cast<float>(threshold_.load(std::memory_order_relaxed));
177
178 // The STFT invokes the callback once per PROCESSED channel per hop
179 // (in channel order) - and it processes min(buffer, prepared)
180 // channels. Modulo by that same effective count, reset per block, or
181 // a narrow buffer over a wider spec would rotate the per-channel
182 // gain memories between hops (channel 0 alternating onto channel 1's
183 // release state - measured 0.004 divergence versus a mono-prepared
184 // twin before this fix).
185 const int nChEff = std::max(1, std::min(buffer.getNumChannels(), numChannels_));
186 callCounter_ = 0;
187
188 // Over-subtraction: the learned noise power is scaled by threshold^2.
189 const float overSub = thresh * thresh;
190
191 stft_.processBlock(buffer, [this, learning, floorGain, overSub, nChEff](T* bins, int numBins)
192 {
193 auto& clean = cleanPower_[static_cast<size_t>(callCounter_ % nChEff)];
194 ++callCounter_;
195
196 // The profile is the running MEAN noise power per bin (the
197 // cumulative average over every learned frame, then an
198 // exponential average once kMaxLearnFrames are in).
199 float learnRate = 0.0f;
200 if (learning)
201 {
202 learnFrames_ = std::min(learnFrames_ + 1, kMaxLearnFrames);
203 learnRate = 1.0f / static_cast<float>(learnFrames_);
204 }
205
206 for (int k = 0; k < numBins; ++k)
207 {
208 const float re = static_cast<float>(bins[2 * k]);
209 const float im = static_cast<float>(bins[2 * k + 1]);
210 const float power = re * re + im * im;
211
212 auto& noise = profile_[static_cast<size_t>(k)];
213 if (learning)
214 noise += learnRate * (power - noise);
215
216 // Decision-directed Wiener gain (Ephraim-Malah a priori SNR):
217 // the a priori SNR mixes the previous frame's clean estimate
218 // with this frame's excess power, so it follows speech and
219 // music onsets at once while random noise flicker cannot open
220 // a bin (the hard per-bin gate did exactly that: musical noise).
221 float gain = 1.0f;
222 const float lambda = overSub * noise;
223 auto& prevClean = clean[static_cast<size_t>(k)];
224 if (lambda > 0.0f)
225 {
226 const float post = power / lambda;
227 const float prio = std::max(kDDAlpha * prevClean / lambda
228 + (1.0f - kDDAlpha) * std::max(post - 1.0f, 0.0f),
229 kMinPrioriSnr);
230 gain = std::max(prio / (1.0f + prio), floorGain);
231 }
232 prevClean = gain * gain * power;
233
234 bins[2 * k] = static_cast<T>(re * gain);
235 bins[2 * k + 1] = static_cast<T>(im * gain);
236 }
237 });
238 }
239
240private:
242 int numChannels_ = 0;
243 int numBins_ = 0;
244 std::atomic<bool> prepared_ { false };
245
246 static constexpr float kDDAlpha = 0.98f;
247 static constexpr float kMinPrioriSnr = 0.003162f;
248 static constexpr int kMaxLearnFrames = 4096;
249
250 std::vector<float> profile_;
251 int learnFrames_ = 0;
252 std::vector<std::vector<float>> cleanPower_;
253 int callCounter_ = 0;
254
255 std::atomic<bool> learning_ { false };
256 std::atomic<T> reduction_ { T(18) };
257 std::atomic<T> threshold_ { T(2) };
258};
259
260} // namespace dspark
Non-owning view over audio channel data.
Definition AudioBuffer.h:50
Learn-a-profile spectral noise reduction (hiss/hum/room-tone).
void setThreshold(T factor) noexcept
Noise over-subtraction factor over the learned profile, in magnitude [1, 8] (default 2): the gain rul...
void setReduction(T db) noexcept
Maximum attenuation of noise bins in dB [0, 40] (default 18): the floor of the per-bin gain....
bool setState(const uint8_t *data, size_t size)
Restores parameters from a blob (tolerant; rejects foreign ids).
T getReduction() const noexcept
T getThreshold() const noexcept
std::vector< uint8_t > getState() const
Serializes the parameter state (the learned profile is material- dependent content,...
void processBlock(AudioBufferView< T > buffer) noexcept
Processes a block in-place. Pass-through until prepare() succeeds.
void clearProfile() noexcept
Forgets the learned noise profile (stream-owner thread).
int getLatency() const noexcept
Latency in samples (the STFT pipeline's).
bool getLearning() const noexcept
void prepare(const AudioSpec &spec, int fftSize=2048)
Prepares the STFT pipeline and the per-channel bin state.
void reset() noexcept
Clears signal state and per-bin gain memories (keeps profile).
void setLearning(bool learning) noexcept
While true, incoming audio trains the noise profile.
High-performance STFT analysis-modification-synthesis pipeline.
Tolerant reader: missing keys yield defaults, unknown keys are skipped.
Definition StateBlob.h:161
float read(const char *key, float defaultValue) const
Reads a float, or defaultValue when the key is absent.
Definition StateBlob.h:204
bool isValid() const noexcept
Definition StateBlob.h:199
uint32_t processorId() const noexcept
Definition StateBlob.h:200
Serializes key/value parameters into a versioned blob.
Definition StateBlob.h:53
std::vector< uint8_t > blob() const
Finalizes and returns the blob.
Definition StateBlob.h:105
void write(const char *key, float value)
Writes a float parameter.
Definition StateBlob.h:71
Main namespace for the DSPark framework.
constexpr uint32_t stateId(const char(&tag)[5]) noexcept
Builds a FOURCC processor id, e.g. dspark::stateId("COMP").
Definition StateBlob.h:651
Describes the audio environment for a DSP processor.
Definition AudioSpec.h:37
constexpr bool isValid() const noexcept
Checks if the specification contains valid, processable parameters.
Definition AudioSpec.h:71
int numChannels
Number of audio channels (e.g., 1 = mono, 2 = stereo).
Definition AudioSpec.h:58