diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/Fft.java b/ts3-client/core/src/main/java/com/ts3client/audio/Fft.java index dbc4a8e..7c4a1d3 100644 --- a/ts3-client/core/src/main/java/com/ts3client/audio/Fft.java +++ b/ts3-client/core/src/main/java/com/ts3client/audio/Fft.java @@ -4,18 +4,18 @@ package com.ts3client.audio; * In-place iterative radix-2 Cooley–Tukey FFT shared by the voice DSP stages. * All arrays must have a power-of-two length. Pure math, no platform dependencies. */ -final class Fft { +public final class Fft { private Fft() { } /** Forward transform (unnormalised). */ - static void forward(double[] re, double[] im) { + public static void forward(double[] re, double[] im) { transform(re, im, false); } /** Inverse transform, normalised by {@code 1/n} so it inverts {@link #forward}. */ - static void inverse(double[] re, double[] im) { + public static void inverse(double[] re, double[] im) { transform(re, im, true); int n = re.length; double scale = 1.0 / n; diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/AudioProcessor.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/AudioProcessor.java new file mode 100644 index 0000000..d63d3cb --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/AudioProcessor.java @@ -0,0 +1,150 @@ +package com.ts3client.audio.processing; + +import com.ts3client.audio.processing.agc2.AdaptiveDigitalConfig; +import com.ts3client.audio.processing.agc2.GainController2; +import com.ts3client.audio.processing.ns.NoiseSuppressor; +import com.ts3client.audio.processing.ns.SuppressionLevel; +import com.ts3client.audio.processing.transients.TransientSuppressor; + +import java.util.concurrent.atomic.AtomicBoolean; + +/** + * The microphone pre-processing chain, set up the way the TeamSpeak 3 client configures + * WebRTC's audio processing module. TS3's capture preprocessor applies, per 10 ms at + * 48 kHz: + *
    + *
  1. noise suppression ("Remove background noise") at {@code denoiser_level} 0–3, run + * on the three-band split. TS3 disables APM's high-pass filter, but APM forces it on + * whenever noise suppression runs, so a 100 Hz high-pass comes with it;
  2. + *
  3. transient suppression ("Typing attenuation"), told about key presses;
  4. + *
  5. automatic gain control. TS3 runs WebRTC's legacy AGC1 (adaptive digital, target + * -9 dBFS, 20 dB compression); this runs its successor, AGC2 adaptive digital.
  6. + *
+ * + *

One deliberate difference: APM only splits into bands for band-wise stages, so with noise + * suppression and AGC1 both off its transient detector reads stale band data and never + * fires. TS3 hides this behind AGC1, which forces a split; with AGC2 it would show, so the + * split is done whenever typing attenuation is on. + * + *

Pure DSP with no platform dependencies. Drive {@link #process} from the capture thread; + * the setters and {@link #keyPressed()} may be called from any thread and take effect at the + * next 10 ms frame. With every stage off, audio passes through untouched. + */ +public final class AudioProcessor { + + public static final int SAMPLE_RATE = 48_000; + public static final int FRAME_SIZE = ThreeBandFilterBank.FULL_BAND_SIZE; + + private volatile boolean noiseSuppression; + private volatile SuppressionLevel suppressionLevel = SuppressionLevel.DB_12; + private volatile boolean transientSuppression; + private volatile boolean gainControl; + private final AtomicBoolean keyPressed = new AtomicBoolean(); + + private final ThreeBandFilterBank filterBank = new ThreeBandFilterBank(); + private HighPassFilter highPassFilter; + private NoiseSuppressor noiseSuppressor; + private TransientSuppressor transientSuppressor; + private GainController2 gainController; + + private final float[] fullBand = new float[FRAME_SIZE]; + private final float[][] bands = new float[ThreeBandFilterBank.NUM_BANDS][ThreeBandFilterBank.SPLIT_BAND_SIZE]; + + public void setNoiseSuppression(boolean enabled) { + this.noiseSuppression = enabled; + } + + public void setSuppressionLevel(SuppressionLevel level) { + this.suppressionLevel = level; + } + + public void setTransientSuppression(boolean enabled) { + this.transientSuppression = enabled; + } + + public void setGainControl(boolean enabled) { + this.gainControl = enabled; + } + + /** Reports a key press anywhere on the system; the transient suppressor only acts while typing. */ + public void keyPressed() { + keyPressed.set(true); + } + + /** Drops all filter state; call when (re)starting capture. */ + public void reset() { + highPassFilter = null; + noiseSuppressor = null; + transientSuppressor = null; + gainController = null; + keyPressed.set(false); + } + + /** + * Processes mono samples in [-1, 1] in place. {@code len} must be a multiple of + * {@link #FRAME_SIZE}. + */ + public void process(float[] buf, int len) { + for (int offset = 0; offset + FRAME_SIZE <= len; offset += FRAME_SIZE) { + processFrame(buf, offset); + } + } + + private void processFrame(float[] buf, int offset) { + applyConfig(); + boolean key = keyPressed.getAndSet(false); + if (noiseSuppressor == null && transientSuppressor == null && gainController == null) return; + + // APM carries samples on the 16-bit scale. + for (int i = 0; i < FRAME_SIZE; i++) { + fullBand[i] = Math.max(-1.f, Math.min(1.f, buf[offset + i])) * 32768.f; + } + + if (highPassFilter != null) { + highPassFilter.process(fullBand, FRAME_SIZE); + } + if (noiseSuppressor != null || transientSuppressor != null) { + filterBank.analysis(fullBand, bands); + } + if (noiseSuppressor != null) { + noiseSuppressor.analyze(bands[0]); + noiseSuppressor.process(bands); + filterBank.synthesis(bands, fullBand); + } + if (transientSuppressor != null) { + transientSuppressor.suppress(fullBand, bands[0], key); + } + if (gainController != null) { + gainController.process(fullBand); + } + + for (int i = 0; i < FRAME_SIZE; i++) { + buf[offset + i] = Math.max(-32768.f, Math.min(32768.f, fullBand[i])) * (1.f / 32768.f); + } + } + + /** Creates or drops stages to match the settings, as APM's {@code ApplyConfig} does. */ + private void applyConfig() { + SuppressionLevel level = suppressionLevel; + if (!noiseSuppression) { + noiseSuppressor = null; + highPassFilter = null; + } else if (noiseSuppressor == null || noiseSuppressor.level() != level) { + // APM builds a fresh suppressor on a level change; the statistics restart too. + noiseSuppressor = new NoiseSuppressor(level, ThreeBandFilterBank.NUM_BANDS); + if (highPassFilter == null) highPassFilter = new HighPassFilter(); + } + + if (!transientSuppression) { + transientSuppressor = null; + } else if (transientSuppressor == null) { + transientSuppressor = new TransientSuppressor(SAMPLE_RATE / ThreeBandFilterBank.NUM_BANDS); + } + + if (!gainControl) { + gainController = null; + } else if (gainController == null) { + gainController = new GainController2(AdaptiveDigitalConfig.DEFAULT); + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/HighPassFilter.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/HighPassFilter.java new file mode 100644 index 0000000..48a38dd --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/HighPassFilter.java @@ -0,0 +1,32 @@ +package com.ts3client.audio.processing; + +/** + * APM's high-pass filter at 48 kHz ({@code high_pass_filter.cc}): one second-order + * Butterworth section at 100 Hz, {@code butter(2, 100/24000, 'high')}, in direct form I. + */ +final class HighPassFilter { + + private static final float B0 = 0.99079f; + private static final float B1 = -1.98157f; + private static final float B2 = 0.99079f; + private static final float A1 = -1.98149f; + private static final float A2 = 0.98166f; + + private float x0, x1, y0, y1; + + void process(float[] y, int len) { + float mx0 = x0, mx1 = x1, my0 = y0, my1 = y1; + for (int k = 0; k < len; k++) { + final float tmp = y[k]; + y[k] = B0 * tmp + B1 * mx0 + B2 * mx1 - A1 * my0 - A2 * my1; + mx1 = mx0; + mx0 = tmp; + my1 = my0; + my0 = y[k]; + } + x0 = mx0; + x1 = mx1; + y0 = my0; + y1 = my1; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ThreeBandFilterBank.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ThreeBandFilterBank.java new file mode 100644 index 0000000..faf8bc6 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ThreeBandFilterBank.java @@ -0,0 +1,152 @@ +package com.ts3client.audio.processing; + +/** + * Splits 10 ms of 48 kHz audio into three 16 kHz bands (0–8, 8–16 and + * 16–24 kHz) and merges them back, ported from WebRTC's + * {@code three_band_filter_bank.cc}. It is how APM feeds the noise suppressor at 48 kHz. + * + *

A cosine-modulated polyphase filter bank: the half-bandwidth low-pass prototype (a + * 48-tap Kaiser-windowed FIR, stored as ten non-zero 4-tap phases) is shifted to the three + * band centres by a DCT. + */ +final class ThreeBandFilterBank { + + static final int NUM_BANDS = 3; + static final int FULL_BAND_SIZE = 480; + static final int SPLIT_BAND_SIZE = FULL_BAND_SIZE / NUM_BANDS; + + private static final int SPARSITY = 4; + private static final int STRIDE_LOG2 = 2; + private static final int STRIDE = 1 << STRIDE_LOG2; + private static final int NUM_ZERO_FILTERS = 2; + private static final int FILTER_SIZE = 4; + private static final int MEMORY_SIZE = FILTER_SIZE * STRIDE - 1; + private static final int NUM_NON_ZERO_FILTERS = SPARSITY * NUM_BANDS - NUM_ZERO_FILTERS; + private static final int SUB_SAMPLING = NUM_BANDS; + private static final int ZERO_FILTER_INDEX_1 = 3; + private static final int ZERO_FILTER_INDEX_2 = 9; + + private static final float[][] FILTER_COEFFS = { + {-0.00047749f, -0.00496888f, +0.16547118f, +0.00425496f}, + {-0.00173287f, -0.01585778f, +0.14989004f, +0.00994113f}, + {-0.00304815f, -0.02536082f, +0.12154542f, +0.01157993f}, + {-0.00346946f, -0.02587886f, +0.04760441f, +0.00607594f}, + {-0.00154717f, -0.01136076f, +0.01387458f, +0.00186353f}, + {+0.00186353f, +0.01387458f, -0.01136076f, -0.00154717f}, + {+0.00607594f, +0.04760441f, -0.02587886f, -0.00346946f}, + {+0.00983212f, +0.08543175f, -0.02982767f, -0.00383509f}, + {+0.00994113f, +0.14989004f, -0.01585778f, -0.00173287f}, + {+0.00425496f, +0.16547118f, -0.00496888f, -0.00047749f}}; + + private static final float[][] DCT_MODULATION = { + {2.f, 2.f, 2.f}, + {1.73205077f, 0.f, -1.73205077f}, + {1.f, -2.f, 1.f}, + {-1.f, 2.f, -1.f}, + {-1.73205077f, 0.f, 1.73205077f}, + {-2.f, -2.f, -2.f}, + {-1.73205077f, 0.f, 1.73205077f}, + {-1.f, 2.f, -1.f}, + {1.f, -2.f, 1.f}, + {1.73205077f, 0.f, -1.73205077f}}; + + private final float[][] stateAnalysis = new float[NUM_NON_ZERO_FILTERS][MEMORY_SIZE]; + private final float[][] stateSynthesis = new float[NUM_NON_ZERO_FILTERS][MEMORY_SIZE]; + + private final float[] inSubsampled = new float[SPLIT_BAND_SIZE]; + private final float[] outSubsampled = new float[SPLIT_BAND_SIZE]; + + /** Splits {@code in} (480 samples) into {@code out[band]} (160 samples each). */ + void analysis(float[] in, float[][] out) { + for (int band = 0; band < NUM_BANDS; band++) { + java.util.Arrays.fill(out[band], 0, SPLIT_BAND_SIZE, 0.f); + } + + for (int downsamplingIndex = 0; downsamplingIndex < SUB_SAMPLING; downsamplingIndex++) { + for (int k = 0; k < SPLIT_BAND_SIZE; k++) { + inSubsampled[k] = in[(SUB_SAMPLING - 1) - downsamplingIndex + SUB_SAMPLING * k]; + } + + for (int inShift = 0; inShift < STRIDE; inShift++) { + int filterIndex = filterIndex(downsamplingIndex + inShift * SUB_SAMPLING); + if (filterIndex < 0) continue; + + filterCore(FILTER_COEFFS[filterIndex], inSubsampled, inShift, outSubsampled, + stateAnalysis[filterIndex]); + + float[] dct = DCT_MODULATION[filterIndex]; + for (int band = 0; band < NUM_BANDS; band++) { + float[] outBand = out[band]; + for (int n = 0; n < SPLIT_BAND_SIZE; n++) { + outBand[n] += dct[band] * outSubsampled[n]; + } + } + } + } + } + + /** Merges {@code in[band]} (160 samples each) back into {@code out} (480 samples). */ + void synthesis(float[][] in, float[] out) { + java.util.Arrays.fill(out, 0, FULL_BAND_SIZE, 0.f); + for (int upsamplingIndex = 0; upsamplingIndex < SUB_SAMPLING; upsamplingIndex++) { + for (int inShift = 0; inShift < STRIDE; inShift++) { + int filterIndex = filterIndex(upsamplingIndex + inShift * SUB_SAMPLING); + if (filterIndex < 0) continue; + + float[] dct = DCT_MODULATION[filterIndex]; + java.util.Arrays.fill(inSubsampled, 0.f); + for (int band = 0; band < NUM_BANDS; band++) { + float[] inBand = in[band]; + for (int n = 0; n < SPLIT_BAND_SIZE; n++) { + inSubsampled[n] += dct[band] * inBand[n]; + } + } + + filterCore(FILTER_COEFFS[filterIndex], inSubsampled, inShift, outSubsampled, + stateSynthesis[filterIndex]); + + final float upsamplingScaling = SUB_SAMPLING; + for (int k = 0; k < SPLIT_BAND_SIZE; k++) { + out[upsamplingIndex + SUB_SAMPLING * k] += upsamplingScaling * outSubsampled[k]; + } + } + } + } + + /** Maps a polyphase index to its stored filter, or -1 for the two all-zero phases. */ + private static int filterIndex(int index) { + if (index == ZERO_FILTER_INDEX_1 || index == ZERO_FILTER_INDEX_2) return -1; + if (index < ZERO_FILTER_INDEX_1) return index; + return index < ZERO_FILTER_INDEX_2 ? index - 1 : index - 2; + } + + /** Filters {@code in} with the sparse (upsampled by {@code STRIDE}) filter, shifted by {@code inShift}. */ + private static void filterCore(float[] filter, float[] in, int inShift, float[] out, float[] state) { + java.util.Arrays.fill(out, 0.f); + for (int k = 0; k < inShift; k++) { + for (int i = 0, j = MEMORY_SIZE + k - inShift; i < FILTER_SIZE; i++, j -= STRIDE) { + out[k] += state[j] * filter[i]; + } + } + + for (int k = inShift, shift = 0; k < FILTER_SIZE * STRIDE; k++, shift++) { + final int loopLimit = Math.min(FILTER_SIZE, 1 + (shift >> STRIDE_LOG2)); + for (int i = 0, j = shift; i < loopLimit; i++, j -= STRIDE) { + out[k] += in[j] * filter[i]; + } + for (int i = loopLimit, j = MEMORY_SIZE + shift - loopLimit * STRIDE; i < FILTER_SIZE; + i++, j -= STRIDE) { + out[k] += state[j] * filter[i]; + } + } + + for (int k = FILTER_SIZE * STRIDE, shift = FILTER_SIZE * STRIDE - inShift; k < SPLIT_BAND_SIZE; + k++, shift++) { + for (int i = 0, j = shift; i < FILTER_SIZE; i++, j -= STRIDE) { + out[k] += in[j] * filter[i]; + } + } + + System.arraycopy(in, SPLIT_BAND_SIZE - MEMORY_SIZE, state, 0, MEMORY_SIZE); + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/AdaptiveDigitalConfig.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/AdaptiveDigitalConfig.java new file mode 100644 index 0000000..6824254 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/AdaptiveDigitalConfig.java @@ -0,0 +1,11 @@ +package com.ts3client.audio.processing.agc2; + +/** + * AGC2's adaptive digital settings, {@code GainController2::AdaptiveDigital} in + * {@code audio_processing.h}. {@link #DEFAULT} holds WebRTC's defaults. + */ +public record AdaptiveDigitalConfig(float headroomDb, float maxGainDb, float initialGainDb, + float maxGainChangeDbPerSecond, float maxOutputNoiseLevelDbfs) { + + public static final AdaptiveDigitalConfig DEFAULT = new AdaptiveDigitalConfig(5.0f, 50.0f, 15.0f, 6.0f, -50.0f); +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/AdaptiveDigitalGainController.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/AdaptiveDigitalGainController.java new file mode 100644 index 0000000..78174aa --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/AdaptiveDigitalGainController.java @@ -0,0 +1,97 @@ +package com.ts3client.audio.processing.agc2; + +import static com.ts3client.audio.processing.agc2.Agc2Common.FRAME_DURATION_MS; +import static com.ts3client.audio.processing.agc2.Agc2Common.LIMITER_THRESHOLD_FOR_AGC_GAIN_DBFS; +import static com.ts3client.audio.processing.agc2.Agc2Common.VAD_CONFIDENCE_THRESHOLD; + +/** + * Picks and applies the adaptive gain, from {@code agc2/adaptive_digital_gain_controller.cc}: + * enough to bring the speech level (plus headroom) up to just below full scale, limited so + * the noise floor stays under {@code maxOutputNoiseLevelDbfs}, and moved at most + * {@code maxGainChangeDbPerSecond}. It only raises the gain after a run of confident speech. + */ +final class AdaptiveDigitalGainController { + + /** What the controller knows about the current frame. */ + record FrameInfo(float speechProbability, float speechLevelDbfs, boolean speechLevelReliable, + float noiseRmsDbfs, float headroomDb, float limiterEnvelopeDbfs) { + } + + private final GainApplier gainApplier; + private final AdaptiveDigitalConfig config; + private final int adjacentSpeechFramesThreshold; + private final float maxGainChangeDbPer10ms; + private int framesToGainIncreaseAllowed; + private float lastGainDb; + + AdaptiveDigitalGainController(AdaptiveDigitalConfig config, int adjacentSpeechFramesThreshold) { + this.config = config; + this.adjacentSpeechFramesThreshold = adjacentSpeechFramesThreshold; + this.gainApplier = new GainApplier(Agc2Common.dbToRatio(config.initialGainDb())); + this.maxGainChangeDbPer10ms = config.maxGainChangeDbPerSecond() * FRAME_DURATION_MS / 1000.0f; + this.framesToGainIncreaseAllowed = adjacentSpeechFramesThreshold; + this.lastGainDb = config.initialGainDb(); + } + + void process(FrameInfo info, float[] frame, int samples) { + final float inputLevelDbfs = info.speechLevelDbfs() + info.headroomDb(); + final float targetGainDb = limitGainByLowConfidence( + limitGainByNoise(computeGainDb(inputLevelDbfs), info.noiseRmsDbfs()), + lastGainDb, info.limiterEnvelopeDbfs(), info.speechLevelReliable()); + + // Only allow the gain to rise after a run of confident speech frames. + boolean firstConfidentSpeechFrame = false; + if (info.speechProbability() < VAD_CONFIDENCE_THRESHOLD) { + framesToGainIncreaseAllowed = adjacentSpeechFramesThreshold; + } else if (framesToGainIncreaseAllowed > 0) { + framesToGainIncreaseAllowed--; + firstConfidentSpeechFrame = framesToGainIncreaseAllowed == 0; + } + final boolean gainIncreaseAllowed = framesToGainIncreaseAllowed == 0; + + float maxGainIncreaseDb = maxGainChangeDbPer10ms; + if (firstConfidentSpeechFrame) { + // Make up for the frames the increase was held back. + maxGainIncreaseDb *= adjacentSpeechFramesThreshold; + } + + float difference = targetGainDb - lastGainDb; + if (!gainIncreaseAllowed) { + difference = Math.min(difference, 0.0f); + } + final float gainChangeThisFrameDb = Math.max(-maxGainChangeDbPer10ms, Math.min(maxGainIncreaseDb, difference)); + + if (gainChangeThisFrameDb != 0.f) { + gainApplier.setGainFactor(Agc2Common.dbToRatio(lastGainDb + gainChangeThisFrameDb)); + } + gainApplier.applyGain(frame, samples); + lastGainDb = lastGainDb + gainChangeThisFrameDb; + } + + /** The gain that puts the input level at {@code -headroomDb}, capped at {@code maxGainDb}. */ + private float computeGainDb(float inputLevelDbfs) { + if (inputLevelDbfs < -(config.headroomDb() + config.maxGainDb())) { + return config.maxGainDb(); + } + if (inputLevelDbfs < -config.headroomDb()) { + return -config.headroomDb() - inputLevelDbfs; + } + return 0.0f; + } + + private float limitGainByNoise(float targetGainDb, float inputNoiseLevelDbfs) { + final float maxAllowedGainDb = config.maxOutputNoiseLevelDbfs() - inputNoiseLevelDbfs; + return Math.min(targetGainDb, Math.max(maxAllowedGainDb, 0.0f)); + } + + /** Until the speech level is trusted, don't push the limiter's envelope above -1 dBFS. */ + private static float limitGainByLowConfidence(float targetGainDb, float lastGainDb, + float limiterAudioLevelDbfs, boolean estimateIsConfident) { + if (estimateIsConfident || limiterAudioLevelDbfs <= LIMITER_THRESHOLD_FOR_AGC_GAIN_DBFS) { + return targetGainDb; + } + final float limiterLevelDbfsBeforeGain = limiterAudioLevelDbfs - lastGainDb; + final float newTargetGainDb = Math.max(LIMITER_THRESHOLD_FOR_AGC_GAIN_DBFS - limiterLevelDbfsBeforeGain, 0.0f); + return Math.min(newTargetGainDb, targetGainDb); + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/Agc2Common.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/Agc2Common.java new file mode 100644 index 0000000..19924f8 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/Agc2Common.java @@ -0,0 +1,41 @@ +package com.ts3client.audio.processing.agc2; + +/** Constants and level conversions shared by AGC2, from {@code agc2/agc2_common.h} and {@code audio_util.h}. */ +final class Agc2Common { + + static final float MIN_FLOAT_S16_VALUE = -32768.0f; + static final float MAX_FLOAT_S16_VALUE = 32767.0f; + static final float MIN_LEVEL_DBFS = -90.31f; + + static final int FRAME_DURATION_MS = 10; + static final int SUB_FRAMES_IN_FRAME = 20; + + /** Target peak for the limiter's envelope when the speech level isn't trusted yet. */ + static final float LIMITER_THRESHOLD_FOR_AGC_GAIN_DBFS = -1.0f; + + static final int VAD_RESET_PERIOD_MS = 1500; + static final float VAD_CONFIDENCE_THRESHOLD = 0.95f; + static final int ADJACENT_SPEECH_FRAMES_THRESHOLD = 12; + + static final float LEVEL_ESTIMATOR_TIME_TO_CONFIDENCE_MS = 400; + static final float LEVEL_ESTIMATOR_LEAK_FACTOR = 1.0f - 1.0f / LEVEL_ESTIMATOR_TIME_TO_CONFIDENCE_MS; + + static final float SATURATION_PROTECTOR_INITIAL_HEADROOM_DB = 20.0f; + static final int SATURATION_PROTECTOR_BUFFER_SIZE = 4; + + private Agc2Common() { + } + + static float dbToRatio(float db) { + return (float) Math.pow(10.0f, db / 20.0f); + } + + /** Level of a 16-bit-scale amplitude in dBFS, floored at the int16 noise floor. */ + static float floatS16ToDbfs(float v) { + final float minDbfs = -90.30899869919436f; + if (v <= 1.0f) { + return minDbfs; + } + return 20.0f * (float) Math.log10(v) + minDbfs; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/GainApplier.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/GainApplier.java new file mode 100644 index 0000000..2bae21b --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/GainApplier.java @@ -0,0 +1,46 @@ +package com.ts3client.audio.processing.agc2; + +/** + * Applies a gain factor, ramping linearly across the frame from the previous factor, from + * {@code agc2/gain_applier.cc}. + */ +final class GainApplier { + + private float lastGainFactor; + private float currentGainFactor; + + GainApplier(float initialGainFactor) { + this.lastGainFactor = initialGainFactor; + this.currentGainFactor = initialGainFactor; + } + + void setGainFactor(float gainFactor) { + currentGainFactor = gainFactor; + } + + void applyGain(float[] frame, int samples) { + final float last = lastGainFactor; + final float target = currentGainFactor; + lastGainFactor = currentGainFactor; + if (last == target && gainCloseToOne(target)) { + return; + } + if (last == target) { + for (int i = 0; i < samples; i++) { + frame[i] *= target; + } + return; + } + final float increment = (target - last) * (1.f / samples); + float gain = last; + for (int i = 0; i < samples; i++) { + frame[i] *= gain; + gain += increment; + } + } + + private static boolean gainCloseToOne(float gainFactor) { + return 1.f - 1.f / Agc2Common.MAX_FLOAT_S16_VALUE <= gainFactor + && gainFactor <= 1.f + 1.f / Agc2Common.MAX_FLOAT_S16_VALUE; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/GainController2.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/GainController2.java new file mode 100644 index 0000000..e3fac82 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/GainController2.java @@ -0,0 +1,78 @@ +package com.ts3client.audio.processing.agc2; + +import static com.ts3client.audio.processing.agc2.Agc2Common.ADJACENT_SPEECH_FRAMES_THRESHOLD; +import static com.ts3client.audio.processing.agc2.Agc2Common.FRAME_DURATION_MS; +import static com.ts3client.audio.processing.agc2.Agc2Common.SATURATION_PROTECTOR_INITIAL_HEADROOM_DB; +import static com.ts3client.audio.processing.agc2.Agc2Common.VAD_RESET_PERIOD_MS; + +import com.ts3client.audio.vad.RnnSpeechDetector; + +/** + * WebRTC's AGC2 with the adaptive digital controller, as APM runs it + * ({@code gain_controller2.cc}): its own RNN VAD, speech and noise level estimators, a + * saturation protector choosing the headroom, the adaptive gain, and the output limiter. + * + *

TS3 uses the older AGC1 in adaptive-digital mode instead. AGC2 is its replacement in + * WebRTC: it measures the speech level only on frames its RNN VAD calls speech, and caps + * the gain so the measured noise floor stays below -50 dBFS. + * + *

Takes 10 ms mono frames at 48 kHz on the 16-bit scale, in place. + */ +public final class GainController2 { + + private static final int FRAME_SIZE = 480; + + private final RnnSpeechDetector vad = new RnnSpeechDetector(48_000); + private final float[] vadFrame = new float[FRAME_SIZE]; + private final int vadResetPeriodFrames = VAD_RESET_PERIOD_MS / FRAME_DURATION_MS; + private int timeToVadReset = vadResetPeriodFrames; + + private final SpeechLevelEstimator speechLevelEstimator; + private final NoiseFloorEstimator noiseLevelEstimator = new NoiseFloorEstimator(); + private final SaturationProtector saturationProtector = + new SaturationProtector(SATURATION_PROTECTOR_INITIAL_HEADROOM_DB, ADJACENT_SPEECH_FRAMES_THRESHOLD); + private final AdaptiveDigitalGainController adaptiveDigitalController; + private final Limiter limiter = new Limiter(); + + public GainController2(AdaptiveDigitalConfig config) { + this.speechLevelEstimator = new SpeechLevelEstimator(config, ADJACENT_SPEECH_FRAMES_THRESHOLD); + this.adaptiveDigitalController = new AdaptiveDigitalGainController(config, ADJACENT_SPEECH_FRAMES_THRESHOLD); + } + + public void process(float[] frame) { + final float speechProbability = analyzeVad(frame); + + float peak = 0.0f; + float rms = 0.0f; + for (int i = 0; i < FRAME_SIZE; i++) { + peak = Math.max(Math.abs(frame[i]), peak); + rms += frame[i] * frame[i]; + } + final float peakDbfs = Agc2Common.floatS16ToDbfs(peak); + final float rmsDbfs = Agc2Common.floatS16ToDbfs((float) Math.sqrt(rms / FRAME_SIZE)); + + final float noiseRmsDbfs = noiseLevelEstimator.analyze(frame, FRAME_SIZE); + speechLevelEstimator.update(rmsDbfs, peakDbfs, speechProbability); + final float speechLevelDbfs = speechLevelEstimator.levelDbfs(); + + saturationProtector.analyze(speechProbability, peakDbfs, speechLevelDbfs); + final float limiterEnvelopeDbfs = Agc2Common.floatS16ToDbfs(limiter.lastAudioLevel()); + + adaptiveDigitalController.process(new AdaptiveDigitalGainController.FrameInfo( + speechProbability, speechLevelDbfs, speechLevelEstimator.isConfident(), + noiseRmsDbfs, saturationProtector.headroomDb(), limiterEnvelopeDbfs), frame, FRAME_SIZE); + limiter.process(frame, FRAME_SIZE); + } + + /** {@code vad_wrapper.cc}: the RNN's state is reset every 1.5 s so it can't latch. */ + private float analyzeVad(float[] frame) { + if (--timeToVadReset <= 0) { + vad.resetNetwork(); + timeToVadReset = vadResetPeriodFrames; + } + for (int i = 0; i < FRAME_SIZE; i++) { + vadFrame[i] = frame[i] * (1.f / 32768.f); + } + return (float) vad.process(vadFrame, FRAME_SIZE); + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/InterpolatedGainCurve.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/InterpolatedGainCurve.java new file mode 100644 index 0000000..ea086ae --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/InterpolatedGainCurve.java @@ -0,0 +1,76 @@ +package com.ts3client.audio.processing.agc2; + +/** + * The limiter's gain curve as a piecewise-linear table, from + * {@code agc2/interpolated_gain_curve.h}: identity below the knee (about -0.8 dBFS), + * then a soft knee into a hard limit at full scale. + */ +final class InterpolatedGainCurve { + + private static final float MAX_INPUT_LEVEL_LINEAR = 36766.300710566735f; + + private static final float[] X = { + 30057.296875f, 30148.986328125f, 30240.67578125f, 30424.052734375f, + 30607.4296875f, 30790.806640625f, 30974.18359375f, 31157.560546875f, + 31340.939453125f, 31524.31640625f, 31707.693359375f, 31891.0703125f, + 32074.447265625f, 32257.82421875f, 32441.201171875f, 32624.580078125f, + 32807.95703125f, 32991.33203125f, 33174.7109375f, 33358.08984375f, + 33541.46484375f, 33724.84375f, 33819.53515625f, 34009.5390625f, + 34200.05859375f, 34389.81640625f, 34674.48828125f, 35054.375f, + 35434.86328125f, 35814.81640625f, 36195.16796875f, 36575.03125f}; + + private static final float[] M = { + -3.515235675877192989e-07f, -1.050251626111275982e-06f, + -2.085213736791047268e-06f, -3.443004743530764244e-06f, + -4.773849468620028347e-06f, -6.077375928725814447e-06f, + -7.353257842623861507e-06f, -8.601219633419532329e-06f, + -9.821013009059242904e-06f, -1.101243378798244521e-05f, + -1.217532644659513608e-05f, -1.330956911260727793e-05f, + -1.441507538402220234e-05f, -1.549179251014720649e-05f, + -1.653970684856176376e-05f, -1.755882840370759368e-05f, + -1.854918446042574942e-05f, -1.951086778717581183e-05f, + -2.044398024736437947e-05f, -2.1348627342376858e-05f, + -2.222496914328075945e-05f, -2.265374678245279938e-05f, + -2.242570917587727308e-05f, -2.220122041762806475e-05f, + -2.19802095671184361e-05f, -2.176260204578284174e-05f, + -2.133731686626560986e-05f, -2.092481918225530535e-05f, + -2.052459603874012828e-05f, -2.013615448959171772e-05f, + -1.975903069251216948e-05f, -1.939277899509761482e-05f}; + + private static final float[] Q = { + 1.010565876960754395f, 1.031631827354431152f, 1.062929749488830566f, + 1.104239225387573242f, 1.144973039627075195f, 1.185109615325927734f, + 1.224629044532775879f, 1.263512492179870605f, 1.301741957664489746f, + 1.339300632476806641f, 1.376173257827758789f, 1.412345528602600098f, + 1.447803974151611328f, 1.482536554336547852f, 1.516532182693481445f, + 1.549780607223510742f, 1.582272171974182129f, 1.613999366760253906f, + 1.644955039024353027f, 1.675132393836975098f, 1.704526185989379883f, + 1.718986630439758301f, 1.711274504661560059f, 1.703639745712280273f, + 1.696081161499023438f, 1.688597679138183594f, 1.673851132392883301f, + 1.659391283988952637f, 1.645209431648254395f, 1.631297469139099121f, + 1.617647409439086914f, 1.604251742362976074f}; + + private InterpolatedGainCurve() { + } + + /** The gain to apply at {@code inputLevel}, a 16-bit-scale peak level. */ + static float lookUpGainToApply(float inputLevel) { + if (inputLevel <= X[0]) { + return 1.0f; + } + if (inputLevel >= MAX_INPUT_LEVEL_LINEAR) { + // Saturating: scale the peak straight down to full scale. + return 32768.f / inputLevel; + } + // Segment whose start is the last point below inputLevel (lower_bound - 1). + int lo = 0; + int hi = X.length; + while (lo < hi) { + int mid = (lo + hi) >>> 1; + if (X[mid] < inputLevel) lo = mid + 1; + else hi = mid; + } + final int index = lo - 1; + return M[index] * inputLevel + Q[index]; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/Limiter.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/Limiter.java new file mode 100644 index 0000000..aec29d1 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/Limiter.java @@ -0,0 +1,91 @@ +package com.ts3client.audio.processing.agc2; + +import static com.ts3client.audio.processing.agc2.Agc2Common.MAX_FLOAT_S16_VALUE; +import static com.ts3client.audio.processing.agc2.Agc2Common.MIN_FLOAT_S16_VALUE; +import static com.ts3client.audio.processing.agc2.Agc2Common.SUB_FRAMES_IN_FRAME; + +/** + * AGC2's output limiter, from {@code agc2/limiter.cc} and {@code fixed_digital_level_estimator.cc}: + * a peak envelope per 1/20 of the frame (instant attack, slow decay) looked up on the + * limiter's gain curve, with the gain interpolated across the sub-frames. + */ +final class Limiter { + + private static final float ATTACK_FILTER_CONSTANT = 0.0f; + private static final float DECAY_FILTER_CONSTANT = 0.9971259f; + private static final float ATTACK_FIRST_SUBFRAME_INTERPOLATION_POWER = 8.0f; + + private final float[] envelope = new float[SUB_FRAMES_IN_FRAME]; + private final float[] scalingFactors = new float[SUB_FRAMES_IN_FRAME + 1]; + private final float[] perSampleScalingFactors = new float[480]; + private float filterStateLevel; + private float lastScalingFactor = 1.f; + + /** The envelope level of the last sub-frame, on the 16-bit scale. */ + float lastAudioLevel() { + return filterStateLevel; + } + + void process(float[] frame, int samples) { + computeLevel(frame, samples); + + scalingFactors[0] = lastScalingFactor; + for (int i = 0; i < SUB_FRAMES_IN_FRAME; i++) { + scalingFactors[i + 1] = InterpolatedGainCurve.lookUpGainToApply(envelope[i]); + } + + computePerSampleSubframeFactors(samples); + for (int j = 0; j < samples; j++) { + frame[j] = Math.max(MIN_FLOAT_S16_VALUE, Math.min(MAX_FLOAT_S16_VALUE, frame[j] * perSampleScalingFactors[j])); + } + lastScalingFactor = scalingFactors[SUB_FRAMES_IN_FRAME]; + } + + private void computeLevel(float[] frame, int samples) { + final int samplesInSubFrame = samples / SUB_FRAMES_IN_FRAME; + java.util.Arrays.fill(envelope, 0.f); + for (int subFrame = 0; subFrame < SUB_FRAMES_IN_FRAME; subFrame++) { + for (int i = 0; i < samplesInSubFrame; i++) { + envelope[subFrame] = Math.max(envelope[subFrame], Math.abs(frame[subFrame * samplesInSubFrame + i])); + } + } + // Look one sub-frame ahead, so the gain is already down when a peak arrives. + for (int subFrame = 0; subFrame < SUB_FRAMES_IN_FRAME - 1; subFrame++) { + if (envelope[subFrame] < envelope[subFrame + 1]) { + envelope[subFrame] = envelope[subFrame + 1]; + } + } + for (int subFrame = 0; subFrame < SUB_FRAMES_IN_FRAME; subFrame++) { + final float value = envelope[subFrame]; + if (value > filterStateLevel) { + envelope[subFrame] = value * (1 - ATTACK_FILTER_CONSTANT) + filterStateLevel * ATTACK_FILTER_CONSTANT; + } else { + envelope[subFrame] = value * (1 - DECAY_FILTER_CONSTANT) + filterStateLevel * DECAY_FILTER_CONSTANT; + } + filterStateLevel = envelope[subFrame]; + } + } + + private void computePerSampleSubframeFactors(int samples) { + final int subframeSize = samples / SUB_FRAMES_IN_FRAME; + final boolean isAttack = scalingFactors[0] > scalingFactors[1]; + if (isAttack) { + // Upstream intends a power-curve fade here, but its i / n is integer division + // and always 0, so the first sub-frame holds the previous factor. Kept as is. + final int n = subframeSize; + for (int i = 0; i < n; i++) { + perSampleScalingFactors[i] = (float) Math.pow(1.f - i / n, ATTACK_FIRST_SUBFRAME_INTERPOLATION_POWER) + * (scalingFactors[0] - scalingFactors[1]) + scalingFactors[1]; + } + } + for (int i = isAttack ? 1 : 0; i < SUB_FRAMES_IN_FRAME; i++) { + final int subframeStart = i * subframeSize; + final float scalingStart = scalingFactors[i]; + final float scalingEnd = scalingFactors[i + 1]; + final float scalingDiff = (scalingEnd - scalingStart) / subframeSize; + for (int j = 0; j < subframeSize; j++) { + perSampleScalingFactors[subframeStart + j] = scalingStart + scalingDiff * j; + } + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/NoiseFloorEstimator.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/NoiseFloorEstimator.java new file mode 100644 index 0000000..71c70d3 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/NoiseFloorEstimator.java @@ -0,0 +1,90 @@ +package com.ts3client.audio.processing.agc2; + +/** + * Estimates the background noise level, from {@code agc2/noise_level_estimator.cc}: the + * minimum frame energy over 5 s periods, rising only slowly between periods, so AGC2 + * can cap its gain before it amplifies the noise floor. + */ +final class NoiseFloorEstimator { + + private static final int FRAMES_PER_SECOND = 100; + private static final int UPDATE_PERIOD_NUM_FRAMES = 500; + + private int sampleRateHz; + private float minNoiseEnergy; + private boolean firstPeriod; + private boolean preliminaryNoiseEnergySet; + private float preliminaryNoiseEnergy; + private float noiseEnergy; + private int counter; + + NoiseFloorEstimator() { + initialize(48000); + } + + /** @return the noise RMS in dBFS */ + float analyze(float[] frame, int samples) { + final int rate = samples * FRAMES_PER_SECOND; + if (rate != sampleRateHz) { + initialize(rate); + } + float frameEnergy = 0.0f; + for (int i = 0; i < samples; i++) { + frameEnergy += frame[i] * frame[i]; + } + if (frameEnergy <= minNoiseEnergy) { + // Ignore frames at or below the int16 noise floor, e.g. digital silence. + return energyToDbfs(noiseEnergy, samples); + } + + if (preliminaryNoiseEnergySet) { + preliminaryNoiseEnergy = Math.min(preliminaryNoiseEnergy, frameEnergy); + } else { + preliminaryNoiseEnergy = frameEnergy; + preliminaryNoiseEnergySet = true; + } + + if (counter == 0) { + // Period over: move towards the new minimum, slowly if it rose. + firstPeriod = false; + noiseEnergy = smooth(noiseEnergy, preliminaryNoiseEnergy); + counter = UPDATE_PERIOD_NUM_FRAMES; + preliminaryNoiseEnergySet = false; + } else if (firstPeriod) { + noiseEnergy = preliminaryNoiseEnergy; + counter--; + } else { + noiseEnergy = Math.min(noiseEnergy, preliminaryNoiseEnergy); + counter--; + } + return energyToDbfs(noiseEnergy, samples); + } + + private void initialize(int rate) { + sampleRateHz = rate; + firstPeriod = true; + preliminaryNoiseEnergySet = false; + // Two LSBs of RMS: anything quieter is treated as silence. + minNoiseEnergy = rate * 2.0f * 2.0f / FRAMES_PER_SECOND; + preliminaryNoiseEnergy = minNoiseEnergy; + noiseEnergy = minNoiseEnergy; + counter = UPDATE_PERIOD_NUM_FRAMES; + } + + private static float smooth(float currentEstimate, float newEstimate) { + final float attack = 0.5f; + if (currentEstimate < newEstimate) { + return attack * newEstimate + (1.0f - attack) * currentEstimate; + } + return newEstimate; + } + + private static float energyToDbfs(float signalEnergy, int numSamples) { + final float rmsSquare = signalEnergy / numSamples; + final float minDbfs = -90.30899869919436f; + if (rmsSquare <= 1.0f) { + return minDbfs; + } + return 10.0f * (float) Math.log10(rmsSquare) + minDbfs; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/SaturationProtector.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/SaturationProtector.java new file mode 100644 index 0000000..af7a760 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/SaturationProtector.java @@ -0,0 +1,122 @@ +package com.ts3client.audio.processing.agc2; + +import static com.ts3client.audio.processing.agc2.Agc2Common.FRAME_DURATION_MS; +import static com.ts3client.audio.processing.agc2.Agc2Common.MIN_LEVEL_DBFS; +import static com.ts3client.audio.processing.agc2.Agc2Common.SATURATION_PROTECTOR_BUFFER_SIZE; +import static com.ts3client.audio.processing.agc2.Agc2Common.VAD_CONFIDENCE_THRESHOLD; + +/** + * Chooses how much headroom to leave above the speech level, from + * {@code agc2/saturation_protector.cc}: it tracks how far the speech peaks, delayed by a + * few 400 ms super-frames, sit above the speech level, and keeps the margin within + * 12–25 dB. + */ +final class SaturationProtector { + + private static final int PEAK_ENVELOPER_SUPER_FRAME_LENGTH_MS = 400; + private static final float MIN_MARGIN_DB = 12.0f; + private static final float MAX_MARGIN_DB = 25.0f; + private static final float ATTACK = 0.9988493699365052f; + private static final float DECAY = 0.9997697679981565f; + + private static final class State { + float headroomDb; + /** Ring buffer of past super-frame peaks. */ + final float[] peakDelayBuffer = new float[SATURATION_PROTECTOR_BUFFER_SIZE]; + int next; + int size; + float maxPeaksDbfs; + int timeSincePushMs; + + void reset(float initialHeadroomDb) { + headroomDb = initialHeadroomDb; + next = 0; + size = 0; + maxPeaksDbfs = MIN_LEVEL_DBFS; + timeSincePushMs = 0; + } + + void copyFrom(State o) { + headroomDb = o.headroomDb; + System.arraycopy(o.peakDelayBuffer, 0, peakDelayBuffer, 0, peakDelayBuffer.length); + next = o.next; + size = o.size; + maxPeaksDbfs = o.maxPeaksDbfs; + timeSincePushMs = o.timeSincePushMs; + } + + void push(float v) { + peakDelayBuffer[next++] = v; + if (next == peakDelayBuffer.length) next = 0; + if (size < peakDelayBuffer.length) size++; + } + + /** The oldest buffered peak, or {@code fallback} while the buffer is empty. */ + float front(float fallback) { + if (size == 0) return fallback; + return peakDelayBuffer[size == peakDelayBuffer.length ? next : 0]; + } + + void update(float peakDbfs, float speechLevelDbfs) { + maxPeaksDbfs = Math.max(maxPeaksDbfs, peakDbfs); + timeSincePushMs += FRAME_DURATION_MS; + if (timeSincePushMs > PEAK_ENVELOPER_SUPER_FRAME_LENGTH_MS) { + push(maxPeaksDbfs); + maxPeaksDbfs = MIN_LEVEL_DBFS; + timeSincePushMs = 0; + } + + final float delayedPeakDbfs = front(maxPeaksDbfs); + final float differenceDb = delayedPeakDbfs - speechLevelDbfs; + if (differenceDb > headroomDb) { + headroomDb = headroomDb * ATTACK + differenceDb * (1.0f - ATTACK); + } else { + headroomDb = headroomDb * DECAY + differenceDb * (1.0f - DECAY); + } + headroomDb = Math.max(MIN_MARGIN_DB, Math.min(MAX_MARGIN_DB, headroomDb)); + } + } + + private final float initialHeadroomDb; + private final int adjacentSpeechFramesThreshold; + private final State preliminary = new State(); + private final State reliable = new State(); + private int numAdjacentSpeechFrames; + private float headroomDb; + + SaturationProtector(float initialHeadroomDb, int adjacentSpeechFramesThreshold) { + this.initialHeadroomDb = initialHeadroomDb; + this.adjacentSpeechFramesThreshold = adjacentSpeechFramesThreshold; + reset(); + } + + float headroomDb() { + return headroomDb; + } + + void analyze(float speechProbability, float peakDbfs, float speechLevelDbfs) { + if (speechProbability < VAD_CONFIDENCE_THRESHOLD) { + if (adjacentSpeechFramesThreshold > 1) { + if (numAdjacentSpeechFrames >= adjacentSpeechFramesThreshold) { + reliable.copyFrom(preliminary); + } else if (numAdjacentSpeechFrames > 0) { + preliminary.copyFrom(reliable); + } + } + numAdjacentSpeechFrames = 0; + } else { + numAdjacentSpeechFrames++; + preliminary.update(peakDbfs, speechLevelDbfs); + if (numAdjacentSpeechFrames >= adjacentSpeechFramesThreshold) { + headroomDb = preliminary.headroomDb; + } + } + } + + void reset() { + numAdjacentSpeechFrames = 0; + headroomDb = initialHeadroomDb; + preliminary.reset(initialHeadroomDb); + reliable.reset(initialHeadroomDb); + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/SpeechLevelEstimator.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/SpeechLevelEstimator.java new file mode 100644 index 0000000..ca51772 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/agc2/SpeechLevelEstimator.java @@ -0,0 +1,106 @@ +package com.ts3client.audio.processing.agc2; + +import static com.ts3client.audio.processing.agc2.Agc2Common.FRAME_DURATION_MS; +import static com.ts3client.audio.processing.agc2.Agc2Common.LEVEL_ESTIMATOR_LEAK_FACTOR; +import static com.ts3client.audio.processing.agc2.Agc2Common.LEVEL_ESTIMATOR_TIME_TO_CONFIDENCE_MS; +import static com.ts3client.audio.processing.agc2.Agc2Common.SATURATION_PROTECTOR_INITIAL_HEADROOM_DB; +import static com.ts3client.audio.processing.agc2.Agc2Common.VAD_CONFIDENCE_THRESHOLD; + +/** + * Estimates the speech RMS level, from {@code agc2/speech_level_estimator.cc}: a + * speech-probability-weighted average over confidently voiced frames, which only takes + * effect after a run of adjacent speech frames so short noises don't move it. + */ +final class SpeechLevelEstimator { + + /** A level estimate kept as a weighted ratio, and the time left until it is trusted. */ + private static final class State { + int timeToConfidenceMs; + float numerator; + float denominator; + + void copyFrom(State other) { + timeToConfidenceMs = other.timeToConfidenceMs; + numerator = other.numerator; + denominator = other.denominator; + } + } + + private final float initialSpeechLevelDbfs; + private final int adjacentSpeechFramesThreshold; + private final State preliminary = new State(); + private final State reliable = new State(); + private float levelDbfs; + private boolean confident; + private int numAdjacentSpeechFrames; + + SpeechLevelEstimator(AdaptiveDigitalConfig config, int adjacentSpeechFramesThreshold) { + this.initialSpeechLevelDbfs = clamp(-SATURATION_PROTECTOR_INITIAL_HEADROOM_DB + - config.initialGainDb() - config.headroomDb()); + this.adjacentSpeechFramesThreshold = adjacentSpeechFramesThreshold; + reset(); + } + + float levelDbfs() { + return levelDbfs; + } + + boolean isConfident() { + return confident; + } + + void update(float rmsDbfs, float peakDbfs, float speechProbability) { + if (speechProbability < VAD_CONFIDENCE_THRESHOLD) { + // Not a speech frame: commit or roll back the preliminary estimate. + if (adjacentSpeechFramesThreshold > 1) { + if (numAdjacentSpeechFrames >= adjacentSpeechFramesThreshold) { + reliable.copyFrom(preliminary); + } else if (numAdjacentSpeechFrames > 0) { + preliminary.copyFrom(reliable); + } + } + numAdjacentSpeechFrames = 0; + } else { + numAdjacentSpeechFrames++; + final boolean bufferIsFull = preliminary.timeToConfidenceMs == 0; + if (!bufferIsFull) { + preliminary.timeToConfidenceMs -= FRAME_DURATION_MS; + } + final float leakFactor = bufferIsFull ? LEVEL_ESTIMATOR_LEAK_FACTOR : 1.0f; + preliminary.numerator = preliminary.numerator * leakFactor + rmsDbfs * speechProbability; + preliminary.denominator = preliminary.denominator * leakFactor + speechProbability; + final float level = preliminary.numerator / preliminary.denominator; + if (numAdjacentSpeechFrames >= adjacentSpeechFramesThreshold) { + levelDbfs = clamp(level); + } + } + updateIsConfident(); + } + + private void updateIsConfident() { + if (adjacentSpeechFramesThreshold == 1) { + confident = preliminary.timeToConfidenceMs == 0; + return; + } + confident = reliable.timeToConfidenceMs == 0 + || (numAdjacentSpeechFrames >= adjacentSpeechFramesThreshold + && preliminary.timeToConfidenceMs == 0); + } + + void reset() { + resetState(preliminary); + resetState(reliable); + levelDbfs = initialSpeechLevelDbfs; + numAdjacentSpeechFrames = 0; + } + + private void resetState(State state) { + state.timeToConfidenceMs = (int) LEVEL_ESTIMATOR_TIME_TO_CONFIDENCE_MS; + state.numerator = initialSpeechLevelDbfs; + state.denominator = 1.0f; + } + + private static float clamp(float levelDbfs) { + return Math.max(-90.0f, Math.min(30.0f, levelDbfs)); + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/FastMath.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/FastMath.java new file mode 100644 index 0000000..ffdeb47 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/FastMath.java @@ -0,0 +1,49 @@ +package com.ts3client.audio.processing.ns; + +/** + * The noise suppressor's approximate log/exp, ported from {@code ns/fast_math.cc}. They are + * deliberately inexact (the log reads the float's exponent bits), and the suppressor's + * statistics were tuned with these errors in place, so they are kept as they are. + */ +final class FastMath { + + private static final float LOG_OF_2 = 0.69314718056f; + private static final float LOG10_OF_E = 0.4342944819f; + + private FastMath() { + } + + /** Reads the float's bits as an integer: the exponent lands in the top, scaled. */ + private static float fastLog2(float in) { + float out = Float.floatToRawIntBits(in); + out *= 1.1920929e-7f; // 1/2^23 + out -= 126.942695f; // exponent bias + return out; + } + + static float sqrt(float f) { + return (float) Math.sqrt(f); + } + + static float pow2(float p) { + return (float) Math.pow(2.0f, p); + } + + static float pow(float x, float p) { + return pow2(p * fastLog2(x)); + } + + static float log(float x) { + return fastLog2(x) * LOG_OF_2; + } + + static void log(float[] x, float[] y) { + for (int k = 0; k < x.length; k++) { + y[k] = log(x[k]); + } + } + + static float exp(float x) { + return pow(10.f, x * LOG10_OF_E); + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NoiseEstimator.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NoiseEstimator.java new file mode 100644 index 0000000..5a3eec2 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NoiseEstimator.java @@ -0,0 +1,145 @@ +package com.ts3client.audio.processing.ns; + +import static com.ts3client.audio.processing.ns.NsCommon.FFT_SIZE_BY_2_PLUS_1; +import static com.ts3client.audio.processing.ns.NsCommon.SHORT_STARTUP_PHASE_BLOCKS; + +/** + * The noise spectrum, from {@code ns/noise_estimator.cc}: a quantile estimate, blended during + * the first half second with a parametric white/pink model fitted to the input, then + * refined each block according to how likely each bin holds speech. + */ +final class NoiseEstimator { + + private static final int BINS = FFT_SIZE_BY_2_PLUS_1; + private static final int START_BAND = 5; + + /** Natural log of the bin index; the first five are unused. */ + private static final float[] LOG_TABLE = { + 0.f, 0.f, 0.f, 0.f, 0.f, 1.609438f, 1.791759f, + 1.945910f, 2.079442f, 2.197225f, 2.302585f, 2.397895f, 2.484907f, 2.564949f, + 2.639057f, 2.708050f, 2.772589f, 2.833213f, 2.890372f, 2.944439f, 2.995732f, + 3.044522f, 3.091043f, 3.135494f, 3.178054f, 3.218876f, 3.258097f, 3.295837f, + 3.332205f, 3.367296f, 3.401197f, 3.433987f, 3.465736f, 3.496507f, 3.526361f, + 3.555348f, 3.583519f, 3.610918f, 3.637586f, 3.663562f, 3.688879f, 3.713572f, + 3.737669f, 3.761200f, 3.784190f, 3.806663f, 3.828641f, 3.850147f, 3.871201f, + 3.891820f, 3.912023f, 3.931826f, 3.951244f, 3.970292f, 3.988984f, 4.007333f, + 4.025352f, 4.043051f, 4.060443f, 4.077538f, 4.094345f, 4.110874f, 4.127134f, + 4.143135f, 4.158883f, 4.174387f, 4.189655f, 4.204693f, 4.219508f, 4.234107f, + 4.248495f, 4.262680f, 4.276666f, 4.290460f, 4.304065f, 4.317488f, 4.330733f, + 4.343805f, 4.356709f, 4.369448f, 4.382027f, 4.394449f, 4.406719f, 4.418841f, + 4.430817f, 4.442651f, 4.454347f, 4.465908f, 4.477337f, 4.488636f, 4.499810f, + 4.510859f, 4.521789f, 4.532599f, 4.543295f, 4.553877f, 4.564348f, 4.574711f, + 4.584968f, 4.595119f, 4.605170f, 4.615121f, 4.624973f, 4.634729f, 4.644391f, + 4.653960f, 4.663439f, 4.672829f, 4.682131f, 4.691348f, 4.700480f, 4.709530f, + 4.718499f, 4.727388f, 4.736198f, 4.744932f, 4.753591f, 4.762174f, 4.770685f, + 4.779124f, 4.787492f, 4.795791f, 4.804021f, 4.812184f, 4.820282f, 4.828314f, + 4.836282f, 4.844187f, 4.852030f}; + + private final SuppressionLevel params; + private final QuantileNoiseEstimator quantileNoiseEstimator = new QuantileNoiseEstimator(); + + private float whiteNoiseLevel; + private float pinkNoiseNumerator; + private float pinkNoiseExp; + final float[] prevNoiseSpectrum = new float[BINS]; + final float[] conservativeNoiseSpectrum = new float[BINS]; + final float[] parametricNoiseSpectrum = new float[BINS]; + final float[] noiseSpectrum = new float[BINS]; + + NoiseEstimator(SuppressionLevel params) { + this.params = params; + } + + void prepareAnalysis() { + System.arraycopy(noiseSpectrum, 0, prevNoiseSpectrum, 0, BINS); + } + + /** First-stage estimate for this block, before the speech probability is known. */ + void preUpdate(int numAnalyzedFrames, float[] signalSpectrum, float signalSpectralSum) { + quantileNoiseEstimator.estimate(signalSpectrum, noiseSpectrum); + if (numAnalyzedFrames >= SHORT_STARTUP_PHASE_BLOCKS) return; + + // Fit white and pink noise models to the spectrum while the quantiles settle. + float sumLogILogMagn = 0.f; + float sumLogI = 0.f; + float sumLogISquare = 0.f; + float sumLogMagn = 0.f; + for (int i = START_BAND; i < BINS; i++) { + float logI = LOG_TABLE[i]; + sumLogI += logI; + sumLogISquare += logI * logI; + float logSignal = FastMath.log(signalSpectrum[i]); + sumLogMagn += logSignal; + sumLogILogMagn += logI * logSignal; + } + + final float oneByBins = 1.f / BINS; + whiteNoiseLevel += signalSpectralSum * oneByBins * params.overSubtractionFactor; + + float denom = sumLogISquare * (BINS - START_BAND) - sumLogI * sumLogI; + float num = sumLogISquare * sumLogMagn - sumLogI * sumLogILogMagn; + float pinkNoiseAdjustment = num / denom; + pinkNoiseAdjustment = Math.max(pinkNoiseAdjustment, 0.f); + pinkNoiseNumerator += pinkNoiseAdjustment; + num = sumLogI * sumLogMagn - (BINS - START_BAND) * sumLogILogMagn; + pinkNoiseAdjustment = num / denom; + pinkNoiseAdjustment = Math.max(Math.min(pinkNoiseAdjustment, 1.f), 0.f); + pinkNoiseExp += pinkNoiseAdjustment; + + final float oneByNumAnalyzedFramesPlus1 = 1.f / (numAnalyzedFrames + 1.f); + float parametricExp = 0.f; + float parametricNum = 0.f; + if (pinkNoiseExp > 0.f) { + parametricNum = FastMath.exp(pinkNoiseNumerator * oneByNumAnalyzedFramesPlus1); + parametricNum *= numAnalyzedFrames + 1.f; + parametricExp = pinkNoiseExp * oneByNumAnalyzedFramesPlus1; + } + + for (int i = 0; i < BINS; i++) { + if (pinkNoiseExp == 0.f) { + parametricNoiseSpectrum[i] = whiteNoiseLevel; + } else { + float useBand = i < START_BAND ? START_BAND : i; + parametricNoiseSpectrum[i] = parametricNum / FastMath.pow(useBand, parametricExp); + } + } + + final float oneByShortStartupPhaseBlocks = 1.f / SHORT_STARTUP_PHASE_BLOCKS; + for (int i = 0; i < BINS; i++) { + noiseSpectrum[i] *= numAnalyzedFrames; + float tmp = parametricNoiseSpectrum[i] * (SHORT_STARTUP_PHASE_BLOCKS - numAnalyzedFrames); + noiseSpectrum[i] += tmp * oneByNumAnalyzedFramesPlus1; + noiseSpectrum[i] *= oneByShortStartupPhaseBlocks; + } + } + + /** Refines the estimate, updating slowly where speech is likely. */ + void postUpdate(float[] speechProbability, float[] signalSpectrum) { + final float noiseUpdate = 0.9f; + final float probRange = .2f; + float gamma = noiseUpdate; + for (int i = 0; i < BINS; i++) { + final float probSpeech = speechProbability[i]; + final float probNonSpeech = 1.f - probSpeech; + + float noiseUpdateTmp = gamma * prevNoiseSpectrum[i] + + (1.f - gamma) * (probNonSpeech * signalSpectrum[i] + probSpeech * prevNoiseSpectrum[i]); + + float gammaOld = gamma; + gamma = probSpeech > probRange ? .99f : noiseUpdate; + + if (probSpeech < probRange) { + conservativeNoiseSpectrum[i] += 0.05f * (signalSpectrum[i] - conservativeNoiseSpectrum[i]); + } + + if (gamma == gammaOld) { + noiseSpectrum[i] = noiseUpdateTmp; + } else { + noiseSpectrum[i] = gamma * prevNoiseSpectrum[i] + + (1.f - gamma) * (probNonSpeech * signalSpectrum[i] + probSpeech * prevNoiseSpectrum[i]); + // A downward update is always safe. + noiseSpectrum[i] = Math.min(noiseSpectrum[i], noiseUpdateTmp); + } + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NoiseSuppressor.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NoiseSuppressor.java new file mode 100644 index 0000000..f782d26 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NoiseSuppressor.java @@ -0,0 +1,320 @@ +package com.ts3client.audio.processing.ns; + +import static com.ts3client.audio.processing.ns.NsCommon.FFT_SIZE; +import static com.ts3client.audio.processing.ns.NsCommon.FFT_SIZE_BY_2_PLUS_1; +import static com.ts3client.audio.processing.ns.NsCommon.NS_FRAME_SIZE; +import static com.ts3client.audio.processing.ns.NsCommon.OVERLAP_SIZE; + +import com.ts3client.audio.Fft; + +/** + * WebRTC's noise suppressor ({@code ns/noise_suppressor.cc}), the module behind TS3's + * "Remove background noise", for one channel. + * + *

It works on band-split audio in 10 ms frames, samples on the 16-bit scale: the + * lowest band (0–8 kHz at 16 kHz) is filtered per bin, and the upper bands, if + * any, are delayed to match and scaled by one gain derived from the top of the lowest band. + * As in APM, call {@link #analyze} and then {@link #process} on every frame. + */ +public final class NoiseSuppressor { + + private static final int BINS = FFT_SIZE_BY_2_PLUS_1; + + /** Rising half of the hybrid Hann/flat analysis and synthesis window. */ + private static final float[] BLOCKS_160W256_FIRST_HALF = { + 0.00000000f, 0.01636173f, 0.03271908f, 0.04906767f, 0.06540313f, + 0.08172107f, 0.09801714f, 0.11428696f, 0.13052619f, 0.14673047f, + 0.16289547f, 0.17901686f, 0.19509032f, 0.21111155f, 0.22707626f, + 0.24298018f, 0.25881905f, 0.27458862f, 0.29028468f, 0.30590302f, + 0.32143947f, 0.33688985f, 0.35225005f, 0.36751594f, 0.38268343f, + 0.39774847f, 0.41270703f, 0.42755509f, 0.44228869f, 0.45690388f, + 0.47139674f, 0.48576339f, 0.50000000f, 0.51410274f, 0.52806785f, + 0.54189158f, 0.55557023f, 0.56910015f, 0.58247770f, 0.59569930f, + 0.60876143f, 0.62166057f, 0.63439328f, 0.64695615f, 0.65934582f, + 0.67155895f, 0.68359230f, 0.69544264f, 0.70710678f, 0.71858162f, + 0.72986407f, 0.74095113f, 0.75183981f, 0.76252720f, 0.77301045f, + 0.78328675f, 0.79335334f, 0.80320753f, 0.81284668f, 0.82226822f, + 0.83146961f, 0.84044840f, 0.84920218f, 0.85772861f, 0.86602540f, + 0.87409034f, 0.88192126f, 0.88951608f, 0.89687274f, 0.90398929f, + 0.91086382f, 0.91749450f, 0.92387953f, 0.93001722f, 0.93590593f, + 0.94154407f, 0.94693013f, 0.95206268f, 0.95694034f, 0.96156180f, + 0.96592583f, 0.97003125f, 0.97387698f, 0.97746197f, 0.98078528f, + 0.98384601f, 0.98664333f, 0.98917651f, 0.99144486f, 0.99344778f, + 0.99518473f, 0.99665524f, 0.99785892f, 0.99879546f, 0.99946459f, + 0.99986614f}; + + private final int numBands; + private final SuppressionLevel params; + private int numAnalyzedFrames = -1; + + private final SpeechProbabilityEstimator speechProbabilityEstimator = new SpeechProbabilityEstimator(); + private final WienerFilter wienerFilter; + private final NoiseEstimator noiseEstimator; + private final float[] prevAnalysisSignalSpectrum = new float[BINS]; + private final float[] analyzeAnalysisMemory = new float[OVERLAP_SIZE]; + private final float[] processAnalysisMemory = new float[OVERLAP_SIZE]; + private final float[] processSynthesisMemory = new float[OVERLAP_SIZE]; + private final float[][] processDelayMemory; + + // Scratch. + private final float[] extendedFrame = new float[FFT_SIZE]; + private final float[] real = new float[FFT_SIZE]; + private final float[] imag = new float[FFT_SIZE]; + private final double[] fftRe = new double[FFT_SIZE]; + private final double[] fftIm = new double[FFT_SIZE]; + private final float[] signalSpectrum = new float[BINS]; + private final float[] priorSnr = new float[BINS]; + private final float[] postSnr = new float[BINS]; + private final float[] delayedFrame = new float[NS_FRAME_SIZE]; + + /** + * @param level how hard to cut + * @param numBands 1 for 16 kHz audio, 3 for 48 kHz split by the three-band filter bank + */ + public NoiseSuppressor(SuppressionLevel level, int numBands) { + this.params = level; + this.numBands = numBands; + this.wienerFilter = new WienerFilter(level); + this.noiseEstimator = new NoiseEstimator(level); + java.util.Arrays.fill(prevAnalysisSignalSpectrum, 1.f); + this.processDelayMemory = new float[Math.max(0, numBands - 1)][OVERLAP_SIZE]; + } + + public SuppressionLevel level() { + return params; + } + + /** Updates the noise and speech statistics from the lowest band of this frame. */ + public void analyze(float[] band0) { + noiseEstimator.prepareAnalysis(); + + // Leave the statistics alone on digital silence: learning from zeros would make any + // later signal look like speech and nothing would be suppressed for a long while. + if (energyOfExtendedFrame(band0, analyzeAnalysisMemory) <= 0.f) { + return; + } + + if (++numAnalyzedFrames < 0) { + numAnalyzedFrames = 0; + } + + formExtendedFrame(band0, analyzeAnalysisMemory, extendedFrame); + applyFilterBankWindow(extendedFrame); + fft(extendedFrame, real, imag); + computeMagnitudeSpectrum(real, imag, signalSpectrum); + + float signalEnergy = 0.f; + for (int i = 0; i < BINS; i++) { + signalEnergy += real[i] * real[i] + imag[i] * imag[i]; + } + signalEnergy /= BINS; + + float signalSpectralSum = 0.f; + for (int i = 0; i < BINS; i++) { + signalSpectralSum += signalSpectrum[i]; + } + + noiseEstimator.preUpdate(numAnalyzedFrames, signalSpectrum, signalSpectralSum); + computeSnr(wienerFilter.filter, prevAnalysisSignalSpectrum, signalSpectrum, + noiseEstimator.prevNoiseSpectrum, noiseEstimator.noiseSpectrum, priorSnr, postSnr); + speechProbabilityEstimator.update(numAnalyzedFrames, priorSnr, postSnr, + noiseEstimator.conservativeNoiseSpectrum, signalSpectrum, signalSpectralSum, signalEnergy); + noiseEstimator.postUpdate(speechProbabilityEstimator.speechProbability, signalSpectrum); + + System.arraycopy(signalSpectrum, 0, prevAnalysisSignalSpectrum, 0, BINS); + } + + /** + * Suppresses noise in place. {@code bands[0]} is the lowest band; any further bands are + * the upper ones. Each holds {@value NsCommon#NS_FRAME_SIZE} samples on the 16-bit scale. + */ + public void process(float[][] bands) { + float[] band0 = bands[0]; + formExtendedFrame(band0, processAnalysisMemory, extendedFrame); + applyFilterBankWindow(extendedFrame); + float energyBeforeFiltering = energyOfExtendedFrame(extendedFrame); + + fft(extendedFrame, real, imag); + computeMagnitudeSpectrum(real, imag, signalSpectrum); + + wienerFilter.update(numAnalyzedFrames, noiseEstimator.noiseSpectrum, noiseEstimator.prevNoiseSpectrum, + noiseEstimator.parametricNoiseSpectrum, signalSpectrum); + + float upperBandGain = 1.f; + if (numBands > 1) { + upperBandGain = computeUpperBandsGain(params.minimumAttenuatingGain, wienerFilter.filter, + speechProbabilityEstimator.speechProbability, prevAnalysisSignalSpectrum, signalSpectrum); + } + + float[] filter = wienerFilter.filter; + for (int i = 0; i < BINS; i++) { + real[i] *= filter[i]; + imag[i] *= filter[i]; + } + ifft(real, imag, extendedFrame); + + float energyAfterFiltering = energyOfExtendedFrame(extendedFrame); + applyFilterBankWindow(extendedFrame); + + float gainAdjustment = wienerFilter.computeOverallScalingFactor(numAnalyzedFrames, + speechProbabilityEstimator.priorSpeechProbability, energyBeforeFiltering, energyAfterFiltering); + for (int i = 0; i < FFT_SIZE; i++) { + extendedFrame[i] = gainAdjustment * extendedFrame[i]; + } + + overlapAndAdd(extendedFrame, processSynthesisMemory, band0); + + for (int b = 1; b < numBands; b++) { + // Delay the upper bands to line up with the lowest band's filter bank. + float[] band = bands[b]; + delaySignal(band, processDelayMemory[b - 1], delayedFrame); + for (int j = 0; j < NS_FRAME_SIZE; j++) { + band[j] = upperBandGain * delayedFrame[j]; + } + } + + for (int b = 0; b < numBands; b++) { + float[] band = bands[b]; + for (int j = 0; j < NS_FRAME_SIZE; j++) { + band[j] = Math.min(Math.max(band[j], -32768.f), 32767.f); + } + } + } + + private static void applyFilterBankWindow(float[] x) { + for (int i = 0; i < 96; i++) { + x[i] = BLOCKS_160W256_FIRST_HALF[i] * x[i]; + } + for (int i = 161, k = 95; i < FFT_SIZE; i++, k--) { + x[i] = BLOCKS_160W256_FIRST_HALF[k] * x[i]; + } + } + + /** Prepends the kept tail of the previous frames, and keeps this frame's tail. */ + private static void formExtendedFrame(float[] frame, float[] oldData, float[] extended) { + System.arraycopy(oldData, 0, extended, 0, OVERLAP_SIZE); + System.arraycopy(frame, 0, extended, OVERLAP_SIZE, NS_FRAME_SIZE); + System.arraycopy(extended, FFT_SIZE - OVERLAP_SIZE, oldData, 0, OVERLAP_SIZE); + } + + private static void overlapAndAdd(float[] extended, float[] overlapMemory, float[] output) { + for (int i = 0; i < OVERLAP_SIZE; i++) { + output[i] = overlapMemory[i] + extended[i]; + } + System.arraycopy(extended, OVERLAP_SIZE, output, OVERLAP_SIZE, NS_FRAME_SIZE - OVERLAP_SIZE); + System.arraycopy(extended, NS_FRAME_SIZE, overlapMemory, 0, OVERLAP_SIZE); + } + + private static void delaySignal(float[] frame, float[] delayBuffer, float[] delayed) { + final int samplesFromFrame = NS_FRAME_SIZE - OVERLAP_SIZE; + System.arraycopy(delayBuffer, 0, delayed, 0, OVERLAP_SIZE); + System.arraycopy(frame, 0, delayed, OVERLAP_SIZE, samplesFromFrame); + System.arraycopy(frame, samplesFromFrame, delayBuffer, 0, OVERLAP_SIZE); + } + + private static float energyOfExtendedFrame(float[] x) { + float energy = 0.f; + for (float v : x) { + energy += v * v; + } + return energy; + } + + private static float energyOfExtendedFrame(float[] frame, float[] oldData) { + float energy = 0.f; + for (float v : oldData) { + energy += v * v; + } + for (int i = 0; i < NS_FRAME_SIZE; i++) { + energy += frame[i] * frame[i]; + } + return energy; + } + + private static void computeMagnitudeSpectrum(float[] real, float[] imag, float[] spectrum) { + spectrum[0] = Math.abs(real[0]) + 1.f; + spectrum[BINS - 1] = Math.abs(real[BINS - 1]) + 1.f; + for (int i = 1; i < BINS - 1; i++) { + spectrum[i] = FastMath.sqrt(real[i] * real[i] + imag[i] * imag[i]) + 1.f; + } + } + + private static void computeSnr(float[] filter, float[] prevSignalSpectrum, float[] signalSpectrum, + float[] prevNoiseSpectrum, float[] noiseSpectrum, + float[] priorSnr, float[] postSnr) { + for (int i = 0; i < BINS; i++) { + float prevEstimate = prevSignalSpectrum[i] / (prevNoiseSpectrum[i] + 0.0001f) * filter[i]; + if (signalSpectrum[i] > noiseSpectrum[i]) { + postSnr[i] = signalSpectrum[i] / (noiseSpectrum[i] + 0.0001f) - 1.f; + } else { + postSnr[i] = 0.f; + } + // Decision-directed estimate of the prior SNR. + priorSnr[i] = 0.98f * prevEstimate + (1.f - 0.98f) * postSnr[i]; + } + } + + /** One gain for the upper bands, from speech probability and gain at the top of band 0. */ + private static float computeUpperBandsGain(float minimumAttenuatingGain, float[] filter, + float[] speechProbability, + float[] prevAnalysisSignalSpectrum, float[] signalSpectrum) { + final int numAvgBins = 32; + final float oneByNumAvgBins = 1.f / numAvgBins; + float avgProbSpeech = 0.f; + float avgFilterGain = 0.f; + for (int i = BINS - numAvgBins - 1; i < BINS - 1; i++) { + avgProbSpeech += speechProbability[i]; + avgFilterGain += filter[i]; + } + avgProbSpeech = avgProbSpeech * oneByNumAvgBins; + avgFilterGain = avgFilterGain * oneByNumAvgBins; + + // Speech removed between analysis and processing (by an AEC) should not count. + float sumAnalysisSpectrum = 0.f; + float sumProcessingSpectrum = 0.f; + for (int i = 0; i < BINS; i++) { + sumAnalysisSpectrum += prevAnalysisSignalSpectrum[i]; + sumProcessingSpectrum += signalSpectrum[i]; + } + avgProbSpeech *= sumProcessingSpectrum / sumAnalysisSpectrum; + + float gain = 0.5f * (1.f + (float) Math.tanh(2.f * avgProbSpeech - 1.f)); + if (avgProbSpeech >= 0.5f) { + gain = 0.25f * gain + 0.75f * avgFilterGain; + } else { + gain = 0.5f * gain + 0.5f * avgFilterGain; + } + return Math.min(Math.max(gain, minimumAttenuatingGain), 1.f); + } + + private void fft(float[] timeData, float[] re, float[] im) { + for (int i = 0; i < FFT_SIZE; i++) { + fftRe[i] = timeData[i]; + fftIm[i] = 0; + } + Fft.forward(fftRe, fftIm); + for (int i = 0; i < BINS; i++) { + re[i] = (float) fftRe[i]; + im[i] = (float) fftIm[i]; + } + im[0] = 0; + im[BINS - 1] = 0; + } + + private void ifft(float[] re, float[] im, float[] timeData) { + fftRe[0] = re[0]; + fftIm[0] = 0; + fftRe[BINS - 1] = re[BINS - 1]; + fftIm[BINS - 1] = 0; + for (int i = 1; i < BINS - 1; i++) { + fftRe[i] = re[i]; + fftIm[i] = im[i]; + fftRe[FFT_SIZE - i] = re[i]; + fftIm[FFT_SIZE - i] = -im[i]; + } + Fft.inverse(fftRe, fftIm); + for (int i = 0; i < FFT_SIZE; i++) { + timeData[i] = (float) fftRe[i]; + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NsCommon.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NsCommon.java new file mode 100644 index 0000000..596eefe --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/NsCommon.java @@ -0,0 +1,29 @@ +package com.ts3client.audio.processing.ns; + +/** + * Sizes and constants of WebRTC's noise suppressor, mirroring {@code ns/ns_common.h}. + * + *

The suppressor works on the lowest 0–8 kHz band at 16 kHz: 10 ms + * frames of 160 samples, analysed with a 256-point FFT over a 96-sample overlap. + */ +final class NsCommon { + + static final int FFT_SIZE = 256; + static final int FFT_SIZE_BY_2_PLUS_1 = FFT_SIZE / 2 + 1; + static final int NS_FRAME_SIZE = 160; + static final int OVERLAP_SIZE = FFT_SIZE - NS_FRAME_SIZE; + + static final int SHORT_STARTUP_PHASE_BLOCKS = 50; + static final int LONG_STARTUP_PHASE_BLOCKS = 200; + static final int FEATURE_UPDATE_WINDOW_SIZE = 500; + + static final float LTR_FEATURE_THR = 0.5f; + static final float BIN_SIZE_LRT = 0.1f; + static final float BIN_SIZE_SPEC_FLAT = 0.05f; + static final float BIN_SIZE_SPEC_DIFF = 0.1f; + + static final int HISTOGRAM_SIZE = 1000; + + private NsCommon() { + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/QuantileNoiseEstimator.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/QuantileNoiseEstimator.java new file mode 100644 index 0000000..ade8ca3 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/QuantileNoiseEstimator.java @@ -0,0 +1,79 @@ +package com.ts3client.audio.processing.ns; + +import static com.ts3client.audio.processing.ns.NsCommon.FFT_SIZE_BY_2_PLUS_1; +import static com.ts3client.audio.processing.ns.NsCommon.LONG_STARTUP_PHASE_BLOCKS; + +import java.util.Arrays; + +/** + * Tracks a low quantile of each bin's log spectrum as the noise floor, from + * {@code ns/quantile_noise_estimator.cc}. Three estimates run staggered so that one of + * them completes a 200-block window every 67 blocks. + */ +final class QuantileNoiseEstimator { + + private static final int SIMULT = 3; + private static final int BINS = FFT_SIZE_BY_2_PLUS_1; + + private final float[] density = new float[SIMULT * BINS]; + private final float[] logQuantile = new float[SIMULT * BINS]; + private final float[] quantile = new float[BINS]; + private final int[] counter = new int[SIMULT]; + private int numUpdates = 1; + + private final float[] logSpectrum = new float[BINS]; + + QuantileNoiseEstimator() { + Arrays.fill(density, 0.3f); + Arrays.fill(logQuantile, 8.f); + final float oneBySimult = 1.f / SIMULT; + for (int i = 0; i < SIMULT; i++) { + counter[i] = (int) Math.floor(LONG_STARTUP_PHASE_BLOCKS * (i + 1.f) * oneBySimult); + } + } + + void estimate(float[] signalSpectrum, float[] noiseSpectrum) { + FastMath.log(signalSpectrum, logSpectrum); + + int quantileIndexToReturn = -1; + for (int s = 0, k = 0; s < SIMULT; s++, k += BINS) { + final float oneByCounterPlus1 = 1.f / (counter[s] + 1.f); + for (int i = 0, j = k; i < BINS; i++, j++) { + final float delta = density[j] > 1.f ? 40.f / density[j] : 40.f; + final float multiplier = delta * oneByCounterPlus1; + if (logSpectrum[i] > logQuantile[j]) { + logQuantile[j] += 0.25f * multiplier; + } else { + logQuantile[j] -= 0.75f * multiplier; + } + + final float width = 0.01f; + final float oneByWidthPlus2 = 1.f / (2.f * width); + if (Math.abs(logSpectrum[i] - logQuantile[j]) < width) { + density[j] = (counter[s] * density[j] + oneByWidthPlus2) * oneByCounterPlus1; + } + } + + if (counter[s] >= LONG_STARTUP_PHASE_BLOCKS) { + counter[s] = 0; + if (numUpdates >= LONG_STARTUP_PHASE_BLOCKS) { + quantileIndexToReturn = k; + } + } + counter[s]++; + } + + // During startup, follow the last estimate so the noise is non-zero from the start. + if (numUpdates < LONG_STARTUP_PHASE_BLOCKS) { + quantileIndexToReturn = BINS * (SIMULT - 1); + numUpdates++; + } + + if (quantileIndexToReturn >= 0) { + for (int i = 0; i < BINS; i++) { + quantile[i] = FastMath.exp(logQuantile[quantileIndexToReturn + i]); + } + } + System.arraycopy(quantile, 0, noiseSpectrum, 0, BINS); + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SignalModelEstimator.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SignalModelEstimator.java new file mode 100644 index 0000000..e5fdfa2 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SignalModelEstimator.java @@ -0,0 +1,255 @@ +package com.ts3client.audio.processing.ns; + +import static com.ts3client.audio.processing.ns.NsCommon.BIN_SIZE_LRT; +import static com.ts3client.audio.processing.ns.NsCommon.BIN_SIZE_SPEC_DIFF; +import static com.ts3client.audio.processing.ns.NsCommon.BIN_SIZE_SPEC_FLAT; +import static com.ts3client.audio.processing.ns.NsCommon.FEATURE_UPDATE_WINDOW_SIZE; +import static com.ts3client.audio.processing.ns.NsCommon.FFT_SIZE_BY_2_PLUS_1; +import static com.ts3client.audio.processing.ns.NsCommon.HISTOGRAM_SIZE; +import static com.ts3client.audio.processing.ns.NsCommon.LTR_FEATURE_THR; + +import java.util.Arrays; + +/** + * The three speech features the suppressor weighs (likelihood ratio, spectral flatness and + * difference from the noise template) and the prior model that decides how to weigh them, + * re-fitted from feature histograms every 500 blocks. Ports {@code signal_model_estimator.cc}, + * {@code prior_signal_model_estimator.cc}, {@code histograms.cc} and the two model structs. + */ +final class SignalModelEstimator { + + private static final int BINS = FFT_SIZE_BY_2_PLUS_1; + private static final float ONE_BY_BINS = 1.f / BINS; + + // SignalModel. + float lrt = LTR_FEATURE_THR; + float spectralFlatness = 0.5f; + float spectralDiff = 0.5f; + final float[] avgLogLrt = new float[BINS]; + + // PriorSignalModel. + float priorLrt = LTR_FEATURE_THR; + float flatnessThreshold = .5f; + float templateDiffThreshold = .5f; + float lrtWeighting = 1.f; + float flatnessWeighting = 0.f; + float differenceWeighting = 0.f; + + private final int[] lrtHistogram = new int[HISTOGRAM_SIZE]; + private final int[] flatnessHistogram = new int[HISTOGRAM_SIZE]; + private final int[] diffHistogram = new int[HISTOGRAM_SIZE]; + + private float diffNormalization; + private float signalEnergySum; + private int histogramAnalysisCounter = 500; + + SignalModelEstimator() { + Arrays.fill(avgLogLrt, LTR_FEATURE_THR); + } + + void adjustNormalization(int numAnalyzedFrames, float signalEnergy) { + diffNormalization *= numAnalyzedFrames; + diffNormalization += signalEnergy; + diffNormalization /= (numAnalyzedFrames + 1); + } + + void update(float[] priorSnr, float[] postSnr, float[] conservativeNoiseSpectrum, + float[] signalSpectrum, float signalSpectralSum, float signalEnergy) { + updateSpectralFlatness(signalSpectrum, signalSpectralSum); + + float diff = computeSpectralDiff(conservativeNoiseSpectrum, signalSpectrum, signalSpectralSum); + spectralDiff += 0.3f * (diff - spectralDiff); + + signalEnergySum += signalEnergy; + + if (--histogramAnalysisCounter > 0) { + updateHistograms(); + } else { + updatePriorModel(); + Arrays.fill(lrtHistogram, 0); + Arrays.fill(flatnessHistogram, 0); + Arrays.fill(diffHistogram, 0); + histogramAnalysisCounter = FEATURE_UPDATE_WINDOW_SIZE; + + signalEnergySum = signalEnergySum / FEATURE_UPDATE_WINDOW_SIZE; + diffNormalization = 0.5f * (signalEnergySum + diffNormalization); + signalEnergySum = 0.f; + } + + updateSpectralLrt(priorSnr, postSnr); + } + + private float computeSpectralDiff(float[] conservativeNoiseSpectrum, float[] signalSpectrum, + float signalSpectralSum) { + float noiseAverage = 0.f; + for (int i = 0; i < BINS; i++) { + noiseAverage += conservativeNoiseSpectrum[i]; + } + noiseAverage = noiseAverage * ONE_BY_BINS; + float signalAverage = signalSpectralSum * ONE_BY_BINS; + + float covariance = 0.f; + float noiseVariance = 0.f; + float signalVariance = 0.f; + for (int i = 0; i < BINS; i++) { + float signalDiff = signalSpectrum[i] - signalAverage; + float noiseDiff = conservativeNoiseSpectrum[i] - noiseAverage; + covariance += signalDiff * noiseDiff; + noiseVariance += noiseDiff * noiseDiff; + signalVariance += signalDiff * signalDiff; + } + covariance *= ONE_BY_BINS; + noiseVariance *= ONE_BY_BINS; + signalVariance *= ONE_BY_BINS; + + float diff = signalVariance - (covariance * covariance) / (noiseVariance + 0.0001f); + return diff / (diffNormalization + 0.0001f); + } + + private void updateSpectralFlatness(float[] signalSpectrum, float signalSpectralSum) { + final float averaging = 0.3f; + for (int i = 1; i < BINS; i++) { + if (signalSpectrum[i] == 0.f) { + spectralFlatness -= averaging * spectralFlatness; + return; + } + } + float num = 0.f; + for (int i = 1; i < BINS; i++) { + num += FastMath.log(signalSpectrum[i]); + } + float denom = signalSpectralSum - signalSpectrum[0]; + denom = denom * ONE_BY_BINS; + num = num * ONE_BY_BINS; + float tmp = FastMath.exp(num) / denom; + spectralFlatness += averaging * (tmp - spectralFlatness); + } + + private void updateSpectralLrt(float[] priorSnr, float[] postSnr) { + for (int i = 0; i < BINS; i++) { + float tmp1 = 1.f + 2.f * priorSnr[i]; + float tmp2 = 2.f * priorSnr[i] / (tmp1 + 0.0001f); + float besselTmp = (postSnr[i] + 1.f) * tmp2; + avgLogLrt[i] += .5f * (besselTmp - FastMath.log(tmp1) - avgLogLrt[i]); + } + float sum = 0.f; + for (int i = 0; i < BINS; i++) { + sum += avgLogLrt[i]; + } + lrt = sum * ONE_BY_BINS; + } + + private void updateHistograms() { + final float oneByBinSizeLrt = 1.f / BIN_SIZE_LRT; + if (lrt < HISTOGRAM_SIZE * BIN_SIZE_LRT && lrt >= 0.f) { + lrtHistogram[(int) (oneByBinSizeLrt * lrt)]++; + } + final float oneByBinSizeSpecFlat = 1.f / BIN_SIZE_SPEC_FLAT; + if (spectralFlatness < HISTOGRAM_SIZE * BIN_SIZE_SPEC_FLAT && spectralFlatness >= 0.f) { + flatnessHistogram[(int) (spectralFlatness * oneByBinSizeSpecFlat)]++; + } + final float oneByBinSizeSpecDiff = 1.f / BIN_SIZE_SPEC_DIFF; + if (spectralDiff < HISTOGRAM_SIZE * BIN_SIZE_SPEC_DIFF && spectralDiff >= 0.f) { + diffHistogram[(int) (spectralDiff * oneByBinSizeSpecDiff)]++; + } + } + + private void updatePriorModel() { + boolean lowLrtFluctuations = updatePriorLrt(); + + float[] flat = firstOfTwoLargestPeaks(BIN_SIZE_SPEC_FLAT, flatnessHistogram); + float flatPeakPosition = flat[0]; + int flatPeakWeight = (int) flat[1]; + float[] diff = firstOfTwoLargestPeaks(BIN_SIZE_SPEC_DIFF, diffHistogram); + float diffPeakPosition = diff[0]; + int diffPeakWeight = (int) diff[1]; + + // Only use the features whose histograms show a clear peak. + final int useSpecFlat = flatPeakWeight < 0.3f * 500 || flatPeakPosition < 0.6f ? 0 : 1; + final int useSpecDiff = diffPeakWeight < 0.3f * 500 || lowLrtFluctuations ? 0 : 1; + + templateDiffThreshold = 1.2f * diffPeakPosition; + templateDiffThreshold = Math.min(1.f, Math.max(0.16f, templateDiffThreshold)); + + float oneByFeatureSum = 1.f / (1.f + useSpecFlat + useSpecDiff); + lrtWeighting = oneByFeatureSum; + + if (useSpecFlat == 1) { + flatnessThreshold = 0.9f * flatPeakPosition; + flatnessThreshold = Math.min(.95f, Math.max(0.1f, flatnessThreshold)); + flatnessWeighting = oneByFeatureSum; + } else { + flatnessWeighting = 0.f; + } + + differenceWeighting = useSpecDiff == 1 ? oneByFeatureSum : 0.f; + } + + /** Re-fits the prior LRT; returns whether the LRT barely fluctuates. */ + private boolean updatePriorLrt() { + float average = 0.f; + float averageCompl = 0.f; + float averageSquared = 0.f; + int count = 0; + for (int i = 0; i < 10; i++) { + float binMid = (i + 0.5f) * BIN_SIZE_LRT; + average += lrtHistogram[i] * binMid; + count += lrtHistogram[i]; + } + if (count > 0) { + average = average / count; + } + for (int i = 0; i < HISTOGRAM_SIZE; i++) { + float binMid = (i + 0.5f) * BIN_SIZE_LRT; + averageSquared += lrtHistogram[i] * binMid * binMid; + averageCompl += lrtHistogram[i] * binMid; + } + final float oneByWindow = 1.f / FEATURE_UPDATE_WINDOW_SIZE; + averageSquared = averageSquared * oneByWindow; + averageCompl = averageCompl * oneByWindow; + + boolean lowLrtFluctuations = averageSquared - average * averageCompl < 0.05f; + final float maxLrt = 1.f; + final float minLrt = .2f; + if (lowLrtFluctuations) { + priorLrt = maxLrt; + } else { + priorLrt = Math.min(maxLrt, Math.max(minLrt, 1.2f * average)); + } + return lowLrtFluctuations; + } + + /** Returns {position, weight} of the larger peak, merged with the runner-up when adjacent. */ + private static float[] firstOfTwoLargestPeaks(float binSize, int[] histogram) { + int peakValue = 0; + int secondaryPeakValue = 0; + float peakPosition = 0.f; + float secondaryPeakPosition = 0.f; + int peakWeight = 0; + int secondaryPeakWeight = 0; + + for (int i = 0; i < HISTOGRAM_SIZE; i++) { + final float binMid = (i + 0.5f) * binSize; + if (histogram[i] > peakValue) { + secondaryPeakValue = peakValue; + secondaryPeakWeight = peakWeight; + secondaryPeakPosition = peakPosition; + + peakValue = histogram[i]; + peakWeight = histogram[i]; + peakPosition = binMid; + } else if (histogram[i] > secondaryPeakValue) { + secondaryPeakValue = histogram[i]; + secondaryPeakWeight = histogram[i]; + secondaryPeakPosition = binMid; + } + } + + if (Math.abs(secondaryPeakPosition - peakPosition) < 2 * binSize + && secondaryPeakWeight > 0.5f * peakWeight) { + peakWeight += secondaryPeakWeight; + peakPosition = 0.5f * (peakPosition + secondaryPeakPosition); + } + return new float[]{peakPosition, peakWeight}; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SpeechProbabilityEstimator.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SpeechProbabilityEstimator.java new file mode 100644 index 0000000..79d24a1 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SpeechProbabilityEstimator.java @@ -0,0 +1,53 @@ +package com.ts3client.audio.processing.ns; + +import static com.ts3client.audio.processing.ns.NsCommon.FFT_SIZE_BY_2_PLUS_1; +import static com.ts3client.audio.processing.ns.NsCommon.LONG_STARTUP_PHASE_BLOCKS; + +/** + * Per-bin speech presence probability, from {@code ns/speech_probability_estimator.cc}: the + * features are mapped through sigmoids around the prior model's thresholds into a prior + * speech probability, which is then combined with each bin's likelihood ratio. + */ +final class SpeechProbabilityEstimator { + + private static final int BINS = FFT_SIZE_BY_2_PLUS_1; + private static final float WIDTH_PRIOR_0 = 4.f; + private static final float WIDTH_PRIOR_1 = 2.f * WIDTH_PRIOR_0; + + private final SignalModelEstimator model = new SignalModelEstimator(); + float priorSpeechProbability = .5f; + final float[] speechProbability = new float[BINS]; + + void update(int numAnalyzedFrames, float[] priorSnr, float[] postSnr, float[] conservativeNoiseSpectrum, + float[] signalSpectrum, float signalSpectralSum, float signalEnergy) { + if (numAnalyzedFrames < LONG_STARTUP_PHASE_BLOCKS) { + model.adjustNormalization(numAnalyzedFrames, signalEnergy); + } + model.update(priorSnr, postSnr, conservativeNoiseSpectrum, signalSpectrum, signalSpectralSum, + signalEnergy); + + float widthPrior = model.lrt < model.priorLrt ? WIDTH_PRIOR_1 : WIDTH_PRIOR_0; + float indicator0 = 0.5f * ((float) Math.tanh(widthPrior * (model.lrt - model.priorLrt)) + 1.f); + + widthPrior = model.spectralFlatness > model.flatnessThreshold ? WIDTH_PRIOR_1 : WIDTH_PRIOR_0; + float indicator1 = 0.5f * ((float) Math.tanh( + 1.f * widthPrior * (model.flatnessThreshold - model.spectralFlatness)) + 1.f); + + widthPrior = model.spectralDiff < model.templateDiffThreshold ? WIDTH_PRIOR_1 : WIDTH_PRIOR_0; + float indicator2 = 0.5f * ((float) Math.tanh( + widthPrior * (model.spectralDiff - model.templateDiffThreshold)) + 1.f); + + float indPrior = model.lrtWeighting * indicator0 + + model.flatnessWeighting * indicator1 + + model.differenceWeighting * indicator2; + + priorSpeechProbability += 0.1f * (indPrior - priorSpeechProbability); + priorSpeechProbability = Math.max(Math.min(priorSpeechProbability, 1.f), 0.01f); + + float gainPrior = (1.f - priorSpeechProbability) / (priorSpeechProbability + 0.0001f); + for (int i = 0; i < BINS; i++) { + float invLrt = FastMath.exp(-model.avgLogLrt[i]); + speechProbability[i] = 1.f / (1.f + gainPrior * invLrt); + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SuppressionLevel.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SuppressionLevel.java new file mode 100644 index 0000000..d0865f6 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/SuppressionLevel.java @@ -0,0 +1,30 @@ +package com.ts3client.audio.processing.ns; + +/** + * How hard the noise suppressor cuts, from {@code ns/suppression_params.cc}. TS3's + * {@code denoiser_level} slider (0–3) picks one of these directly, in this order. + */ +public enum SuppressionLevel { + + DB_6(1.f, 0.5f, false), + DB_12(1.f, 0.25f, true), + DB_18(1.1f, 0.125f, true), + DB_21(1.25f, 0.09f, true); + + final float overSubtractionFactor; + final float minimumAttenuatingGain; + final boolean useAttenuationAdjustment; + + SuppressionLevel(float overSubtractionFactor, float minimumAttenuatingGain, + boolean useAttenuationAdjustment) { + this.overSubtractionFactor = overSubtractionFactor; + this.minimumAttenuatingGain = minimumAttenuatingGain; + this.useAttenuationAdjustment = useAttenuationAdjustment; + } + + /** TS3's {@code denoiser_level}: 0 is 6 dB, 3 is 21 dB; out of range is clamped. */ + public static SuppressionLevel fromDenoiserLevel(int level) { + SuppressionLevel[] all = values(); + return all[Math.max(0, Math.min(all.length - 1, level))]; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/WienerFilter.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/WienerFilter.java new file mode 100644 index 0000000..ef14b1e --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/ns/WienerFilter.java @@ -0,0 +1,93 @@ +package com.ts3client.audio.processing.ns; + +import static com.ts3client.audio.processing.ns.NsCommon.FFT_SIZE_BY_2_PLUS_1; +import static com.ts3client.audio.processing.ns.NsCommon.LONG_STARTUP_PHASE_BLOCKS; +import static com.ts3client.audio.processing.ns.NsCommon.SHORT_STARTUP_PHASE_BLOCKS; + +import java.util.Arrays; + +/** + * The per-bin suppression gain, from {@code ns/wiener_filter.cc}: a Wiener filter on the + * decision-directed prior SNR, floored at the level's minimum gain, plus the overall + * adjustment that eases the cut when it removed a lot of a likely-speech frame. + */ +final class WienerFilter { + + private static final int BINS = FFT_SIZE_BY_2_PLUS_1; + + private final SuppressionLevel params; + private final float[] spectrumPrevProcess = new float[BINS]; + private final float[] initialSpectralEstimate = new float[BINS]; + final float[] filter = new float[BINS]; + + WienerFilter(SuppressionLevel params) { + this.params = params; + Arrays.fill(filter, 1.f); + } + + void update(int numAnalyzedFrames, float[] noiseSpectrum, float[] prevNoiseSpectrum, + float[] parametricNoiseSpectrum, float[] signalSpectrum) { + for (int i = 0; i < BINS; i++) { + // Previous estimate based on the previous frame with the gain filter. + float prevTsa = spectrumPrevProcess[i] / (prevNoiseSpectrum[i] + 0.0001f) * filter[i]; + + float currentTsa; + if (signalSpectrum[i] > noiseSpectrum[i]) { + currentTsa = signalSpectrum[i] / (noiseSpectrum[i] + 0.0001f) - 1.f; + } else { + currentTsa = 0.f; + } + + float snrPrior = 0.98f * prevTsa + (1.f - 0.98f) * currentTsa; + filter[i] = snrPrior / (params.overSubtractionFactor + snrPrior); + filter[i] = Math.max(Math.min(filter[i], 1.f), params.minimumAttenuatingGain); + } + + if (numAnalyzedFrames < SHORT_STARTUP_PHASE_BLOCKS) { + // Blend in the parametric noise model while the estimates settle. + final float oneByShortStartupPhaseBlocks = 1.f / SHORT_STARTUP_PHASE_BLOCKS; + for (int i = 0; i < BINS; i++) { + initialSpectralEstimate[i] += signalSpectrum[i]; + float filterInitial = initialSpectralEstimate[i] + - params.overSubtractionFactor * parametricNoiseSpectrum[i]; + filterInitial /= initialSpectralEstimate[i] + 0.0001f; + filterInitial = Math.max(Math.min(filterInitial, 1.f), params.minimumAttenuatingGain); + + filterInitial *= SHORT_STARTUP_PHASE_BLOCKS - numAnalyzedFrames; + filter[i] *= numAnalyzedFrames; + filter[i] += filterInitial; + filter[i] *= oneByShortStartupPhaseBlocks; + } + } + + System.arraycopy(signalSpectrum, 0, spectrumPrevProcess, 0, BINS); + } + + float computeOverallScalingFactor(int numAnalyzedFrames, float priorSpeechProbability, + float energyBeforeFiltering, float energyAfterFiltering) { + if (!params.useAttenuationAdjustment || numAnalyzedFrames <= LONG_STARTUP_PHASE_BLOCKS) { + return 1.f; + } + + float gain = FastMath.sqrt(energyAfterFiltering / (energyBeforeFiltering + 1.f)); + + final float bLim = 0.5f; + float scaleFactor1 = 1.f; + if (gain > bLim) { + scaleFactor1 = 1.f + 1.3f * (gain - bLim); + if (gain * scaleFactor1 > 1.f) { + scaleFactor1 = 1.f / gain; + } + } + + float scaleFactor2 = 1.f; + if (gain < bLim) { + // Do not reduce the scale too much for pause regions: attenuation is already + // controlled by the gain floor. + gain = Math.max(gain, params.minimumAttenuatingGain); + scaleFactor2 = 1.f - 0.3f * (bLim - gain); + } + + return priorSpeechProbability * scaleFactor1 + (1.f - priorSpeechProbability) * scaleFactor2; + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/TransientDetector.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/TransientDetector.java new file mode 100644 index 0000000..2ace2c5 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/TransientDetector.java @@ -0,0 +1,125 @@ +package com.ts3client.audio.processing.transients; + +import java.util.ArrayDeque; + +/** + * Scores how likely a 10 ms chunk holds a transient such as a key click, from + * {@code transient/transient_detector.cc}. Each wavelet sub-band sample is compared against + * the moving mean and variance of the preceding 30 ms; the summed normalised deviation + * maps through a raised cosine to [0, 1], and the result is the maximum over the last three + * chunks so a click's tail stays covered. + */ +final class TransientDetector { + + private static final int CHUNK_MS = 10; + private static final int TRANSIENT_LENGTH_MS = 30; + private static final int CHUNKS_AT_STARTUP_LEFT_TO_DELETE = TRANSIENT_LENGTH_MS / CHUNK_MS; + private static final float DETECT_THRESHOLD = 16.f; + private static final float PI = 3.14159265358979323846f; + + private final int samplesPerChunk; + private final int leafLength; + private final WaveletPacketTree tree; + private final MovingMoments[] movingMoments = new MovingMoments[WaveletPacketTree.LEAVES]; + private final float[] firstMoments; + private final float[] secondMoments; + private final float[] lastFirstMoment = new float[WaveletPacketTree.LEAVES]; + private final float[] lastSecondMoment = new float[WaveletPacketTree.LEAVES]; + private final ArrayDeque previousResults = new ArrayDeque<>(); + private int chunksAtStartupLeftToDelete = CHUNKS_AT_STARTUP_LEFT_TO_DELETE; + + TransientDetector(int sampleRateHz) { + int perChunk = sampleRateHz * CHUNK_MS / 1000; + int perTransient = sampleRateHz * TRANSIENT_LENGTH_MS / 1000; + perChunk -= perChunk % WaveletPacketTree.LEAVES; + perTransient -= perTransient % WaveletPacketTree.LEAVES; + this.samplesPerChunk = perChunk; + this.tree = new WaveletPacketTree(perChunk); + this.leafLength = tree.leafLength(); + for (int i = 0; i < WaveletPacketTree.LEAVES; i++) { + movingMoments[i] = new MovingMoments(perTransient / WaveletPacketTree.LEAVES); + } + this.firstMoments = new float[leafLength]; + this.secondMoments = new float[leafLength]; + for (int i = 0; i < CHUNKS_AT_STARTUP_LEFT_TO_DELETE; i++) { + previousResults.add(0.f); + } + } + + int samplesPerChunk() { + return samplesPerChunk; + } + + /** Returns a detection score in [0, 1] for the chunk at {@code data[offset..]}. */ + float detect(float[] data, int offset) { + tree.update(data, offset); + + float result = 0.f; + for (int i = 0; i < WaveletPacketTree.LEAVES; i++) { + float[] leaf = tree.leaf(i); + movingMoments[i].calculateMoments(leaf, leafLength, firstMoments, secondMoments); + + // Each sample is judged against the moments of the samples before it. + float unbiased = leaf[0] - lastFirstMoment[i]; + result += unbiased * unbiased / (lastSecondMoment[i] + Float.MIN_NORMAL); + for (int j = 1; j < leafLength; j++) { + unbiased = leaf[j] - firstMoments[j - 1]; + result += unbiased * unbiased / (secondMoments[j - 1] + Float.MIN_NORMAL); + } + + lastFirstMoment[i] = firstMoments[leafLength - 1]; + lastSecondMoment[i] = secondMoments[leafLength - 1]; + } + result /= leafLength; + + // The moments need a full transient length of history before they mean anything. + if (chunksAtStartupLeftToDelete > 0) { + chunksAtStartupLeftToDelete--; + result = 0.f; + } + + if (result >= DETECT_THRESHOLD) { + result = 1.f; + } else { + // Raised cosine from 0 to 1 over [0, threshold), squared. + final float horizontalScaling = PI / DETECT_THRESHOLD; + result = ((float) Math.cos(result * horizontalScaling + PI) + 1.f) * 0.5f; + result *= result; + } + + previousResults.pollFirst(); + previousResults.addLast(result); + float max = 0.f; + for (float r : previousResults) { + max = Math.max(max, r); + } + return max; + } + + /** Running mean and mean square over a fixed window, from {@code moving_moments.cc}. */ + private static final class MovingMoments { + private final int length; + private final float[] queue; + private int head; + private float sum; + private float sumOfSquares; + + MovingMoments(int length) { + this.length = length; + this.queue = new float[length]; + } + + void calculateMoments(float[] in, int n, float[] first, float[] second) { + for (int i = 0; i < n; i++) { + final float oldValue = queue[head]; + queue[head] = in[i]; + head = (head + 1) % length; + + sum += in[i] - oldValue; + sumOfSquares += in[i] * in[i] - oldValue * oldValue; + first[i] = sum / length; + second[i] = Math.max(0.f, sumOfSquares / length); + } + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/TransientSuppressor.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/TransientSuppressor.java new file mode 100644 index 0000000..3260391 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/TransientSuppressor.java @@ -0,0 +1,186 @@ +package com.ts3client.audio.processing.transients; + +import com.ts3client.audio.Fft; + +/** + * WebRTC's transient suppressor ({@code transient/transient_suppressor_impl.cc}), the module + * behind TS3's "Typing attenuation", for one 48 kHz channel. + * + *

It only acts while the user is typing: every key press is reported through + * {@code keyPressed}, the detector switches on with the first press and the suppression + * with the second within about a second; four seconds without a press switch both off. + * While suppressing, spectral peaks that jump above the running spectral mean during a + * detected transient are pulled back towards it. + * + *

TS3 runs it with no analog AGC, so APM hands it a voice probability of 1: the + * suppressor always uses its "soft", voiced restoration, and the unvoiced "hard" path with + * its random phases is never taken. That path is left out here. + * + *

The output is delayed by 544 samples (11.3 ms) whether or not it is suppressing, + * exactly as upstream, so switching on and off doesn't shift the signal. + */ +public final class TransientSuppressor { + + private static final int CHUNK_MS = 10; + private static final float MEAN_IIR_COEFFICIENT = 0.5f; + private static final int MIN_VOICE_BIN = 3; + private static final int MAX_VOICE_BIN = 60; + + private static final int KEYPRESS_PENALTY = 1000 / CHUNK_MS; + private static final int IS_TYPING_THRESHOLD = 1000 / CHUNK_MS; + private static final int CHUNKS_UNTIL_NOT_TYPING = 4000 / CHUNK_MS; + + private static final int DATA_LENGTH = 480; + private static final int ANALYSIS_LENGTH = 1024; + private static final int BUFFER_DELAY = ANALYSIS_LENGTH - DATA_LENGTH; + private static final int COMPLEX_LENGTH = ANALYSIS_LENGTH / 2 + 1; + + /** + * {@code kBlocks480w1024}: a half-sine over [32, 992], zero outside. Upstream stores it + * as a table; this formula reproduces it to within float rounding. + */ + private static final float[] WINDOW = new float[ANALYSIS_LENGTH]; + + static { + for (int i = 32; i < 992; i++) { + WINDOW[i] = (float) Math.sin(Math.PI * (i - 32) / 960.0); + } + } + + private final TransientDetector detector; + private final float[] inBuffer = new float[ANALYSIS_LENGTH]; + private final float[] outBuffer = new float[ANALYSIS_LENGTH]; + private final float[] spectralMean = new float[COMPLEX_LENGTH]; + private final float[] magnitudes = new float[COMPLEX_LENGTH]; + private final float[] meanFactor = new float[COMPLEX_LENGTH]; + private final double[] fftRe = new double[ANALYSIS_LENGTH]; + private final double[] fftIm = new double[ANALYSIS_LENGTH]; + + private float detectorSmoothed; + private int keypressCounter; + private int chunksSinceKeypress; + private boolean detectionEnabled; + private boolean suppressionEnabled; + + /** @param detectionRateHz rate of the detection data passed to {@link #suppress} */ + public TransientSuppressor(int detectionRateHz) { + this.detector = new TransientDetector(detectionRateHz); + // A double sigmoid with its minimum over the voice range (300 Hz - 3 kHz at 16 kHz bins). + final float factorHeight = 10.f; + final float lowSlope = 1.f; + final float highSlope = 0.3f; + for (int i = 0; i < COMPLEX_LENGTH; i++) { + meanFactor[i] = factorHeight / (1.f + (float) Math.exp(lowSlope * (i - MIN_VOICE_BIN))) + + factorHeight / (1.f + (float) Math.exp(highSlope * (MAX_VOICE_BIN - i))); + } + } + + /** + * Processes 10 ms in place. + * + * @param data 480 samples at 48 kHz + * @param detectionData one chunk at the detection rate: APM passes the 0–8 kHz band + * @param keyPressed whether a key went down since the previous call + */ + public void suppress(float[] data, float[] detectionData, boolean keyPressed) { + updateKeypress(keyPressed); + updateBuffers(data); + + if (detectionEnabled) { + float detectorResult = detector.detect(detectionData, 0); + // Rise with the detector, but decay slowly so the ringing of a click is covered. + final float smoothFactor = 0.1f; + detectorSmoothed = detectorResult >= detectorSmoothed + ? detectorResult + : smoothFactor * detectorSmoothed + (1 - smoothFactor) * detectorResult; + suppressChunk(); + } + + // When not suppressing, the input buffer supplies the same delay. + System.arraycopy(suppressionEnabled ? outBuffer : inBuffer, 0, data, 0, DATA_LENGTH); + } + + private void suppressChunk() { + for (int i = 0; i < ANALYSIS_LENGTH; i++) { + fftRe[i] = inBuffer[i] * WINDOW[i]; + fftIm[i] = 0; + } + Fft.forward(fftRe, fftIm); + + for (int i = 0; i < COMPLEX_LENGTH; i++) { + magnitudes[i] = Math.abs((float) fftRe[i]) + Math.abs((float) fftIm[i]); + } + + if (suppressionEnabled) { + softRestoration(); + } + + for (int i = 0; i < COMPLEX_LENGTH; i++) { + spectralMean[i] = (1 - MEAN_IIR_COEFFICIENT) * spectralMean[i] + MEAN_IIR_COEFFICIENT * magnitudes[i]; + } + + // Back to the time domain; the restoration touched the lower half, mirror it. + for (int i = 1; i < COMPLEX_LENGTH - 1; i++) { + fftRe[ANALYSIS_LENGTH - i] = fftRe[i]; + fftIm[ANALYSIS_LENGTH - i] = -fftIm[i]; + } + fftIm[0] = 0; + fftIm[COMPLEX_LENGTH - 1] = 0; + Fft.inverse(fftRe, fftIm); + for (int i = 0; i < ANALYSIS_LENGTH; i++) { + outBuffer[i] += (float) fftRe[i] * WINDOW[i]; + } + } + + /** + * Pulls peaks above the spectral mean back towards it by the detector's confidence, + * skipping those far above the block's voice-band mean: those are more likely voice. + */ + private void softRestoration() { + float blockFrequencyMean = 0; + for (int i = MIN_VOICE_BIN; i < MAX_VOICE_BIN; i++) { + blockFrequencyMean += magnitudes[i]; + } + blockFrequencyMean /= (MAX_VOICE_BIN - MIN_VOICE_BIN); + + for (int i = 0; i < COMPLEX_LENGTH; i++) { + if (magnitudes[i] > spectralMean[i] && magnitudes[i] > 0 + && magnitudes[i] < blockFrequencyMean * meanFactor[i]) { + final float newMagnitude = magnitudes[i] - detectorSmoothed * (magnitudes[i] - spectralMean[i]); + final float ratio = newMagnitude / magnitudes[i]; + fftRe[i] *= ratio; + fftIm[i] *= ratio; + magnitudes[i] = newMagnitude; + } + } + } + + private void updateKeypress(boolean keyPressed) { + if (keyPressed) { + keypressCounter += KEYPRESS_PENALTY; + chunksSinceKeypress = 0; + detectionEnabled = true; + } + keypressCounter = Math.max(0, keypressCounter - 1); + + if (keypressCounter > IS_TYPING_THRESHOLD) { + suppressionEnabled = true; + keypressCounter = 0; + } + + if (detectionEnabled && ++chunksSinceKeypress > CHUNKS_UNTIL_NOT_TYPING) { + detectionEnabled = false; + suppressionEnabled = false; + keypressCounter = 0; + } + } + + private void updateBuffers(float[] data) { + System.arraycopy(inBuffer, DATA_LENGTH, inBuffer, 0, BUFFER_DELAY); + System.arraycopy(data, 0, inBuffer, BUFFER_DELAY, DATA_LENGTH); + if (detectionEnabled) { + System.arraycopy(outBuffer, DATA_LENGTH, outBuffer, 0, BUFFER_DELAY); + java.util.Arrays.fill(outBuffer, BUFFER_DELAY, ANALYSIS_LENGTH, 0.f); + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/WaveletPacketTree.java b/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/WaveletPacketTree.java new file mode 100644 index 0000000..9ea7529 --- /dev/null +++ b/ts3-client/core/src/main/java/com/ts3client/audio/processing/transients/WaveletPacketTree.java @@ -0,0 +1,124 @@ +package com.ts3client.audio.processing.transients; + +/** + * A three-level wavelet packet decomposition with Daubechies-8 filters, ported from + * {@code transient/wpd_tree.cc} and {@code wpd_node.cc}. Each node low- or high-pass + * filters its parent, keeps the odd samples and rectifies them; the eight leaves hand the + * transient detector the signal's envelope in eight sub-bands. + */ +final class WaveletPacketTree { + + static final int LEVELS = 3; + static final int LEAVES = 1 << LEVELS; + + private static final float[] HIGH_PASS = { + -5.44158422430816093862e-02f, 3.12871590914465924627e-01f, + -6.75630736298012846142e-01f, 5.85354683654869090148e-01f, + 1.58291052560238926228e-02f, -2.84015542962428091389e-01f, + -4.72484573997972536787e-04f, 1.28747426620186011803e-01f, + 1.73693010020221083600e-02f, -4.40882539310647192377e-02f, + -1.39810279170155156436e-02f, 8.74609404701565465445e-03f, + 4.87035299301066034600e-03f, -3.91740372995977108837e-04f, + -6.75449405998556772109e-04f, -1.17476784002281916305e-04f}; + + private static final float[] LOW_PASS = { + -1.17476784002281916305e-04f, 6.75449405998556772109e-04f, + -3.91740372995977108837e-04f, -4.87035299301066034600e-03f, + 8.74609404701565465445e-03f, 1.39810279170155156436e-02f, + -4.40882539310647192377e-02f, -1.73693010020221083600e-02f, + 1.28747426620186011803e-01f, 4.72484573997972536787e-04f, + -2.84015542962428091389e-01f, -1.58291052560238926228e-02f, + 5.85354683654869090148e-01f, 6.75630736298012846142e-01f, + 3.12871590914465924627e-01f, 5.44158422430816093862e-02f}; + + /** Heap-ordered: node 1 is the root, node {@code n}'s children are {@code 2n} and {@code 2n+1}. */ + private final Node[] nodes = new Node[(1 << (LEVELS + 1))]; + + WaveletPacketTree(int dataLength) { + nodes[1] = new Node(dataLength, null); + for (int level = 0; level < LEVELS; level++) { + for (int i = 0; i < (1 << level); i++) { + int index = (1 << level) + i; + nodes[2 * index] = new Node(nodes[index].length / 2, LOW_PASS); + nodes[2 * index + 1] = new Node(nodes[index].length / 2, HIGH_PASS); + } + } + } + + void update(float[] data, int offset) { + System.arraycopy(data, offset, nodes[1].data, 0, nodes[1].length); + for (int level = 0; level < LEVELS; level++) { + for (int i = 0; i < (1 << level); i++) { + int index = (1 << level) + i; + nodes[2 * index].update(nodes[index]); + nodes[2 * index + 1].update(nodes[index]); + } + } + } + + /** The data of leaf {@code index}; its first {@link #leafLength()} samples are valid. */ + float[] leaf(int index) { + return nodes[LEAVES + index].data; + } + + int leafLength() { + return nodes[LEAVES].length; + } + + private static final class Node { + final int length; + /** Sized for the parent's length: the filter output lands here before decimation. */ + final float[] data; + private final float[] coefficients; // reversed, as FIRFilterC stores them + private final float[] state; + + Node(int length, float[] coefficients) { + this.length = length; + this.data = new float[2 * length + 1]; + if (coefficients == null) { + this.coefficients = null; + this.state = null; + } else { + int n = coefficients.length; + this.coefficients = new float[n]; + for (int i = 0; i < n; i++) { + this.coefficients[i] = coefficients[n - i - 1]; + } + this.state = new float[n - 1]; + } + } + + void update(Node parent) { + filter(parent.data, parent.length, data); + // Keep the odd samples, then rectify. + for (int i = 0; i < length; i++) { + data[i] = data[2 * i + 1]; + } + for (int i = 0; i < length; i++) { + data[i] = Math.abs(data[i]); + } + } + + /** FIR convolution carrying the last {@code taps - 1} inputs over, as in {@code fir_filter_c.cc}. */ + private void filter(float[] in, int n, float[] out) { + int stateLength = state.length; + for (int i = 0; i < n; i++) { + float acc = 0.f; + int j = 0; + for (; stateLength > i && j < stateLength - i; j++) { + acc += state[i + j] * coefficients[j]; + } + for (; j < coefficients.length; j++) { + acc += in[j + i - stateLength] * coefficients[j]; + } + out[i] = acc; + } + if (n >= stateLength) { + System.arraycopy(in, n - stateLength, state, 0, stateLength); + } else { + System.arraycopy(state, n, state, 0, stateLength - n); + System.arraycopy(in, 0, state, stateLength - n, n); + } + } + } +} diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnSpeechDetector.java b/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnSpeechDetector.java index 4e94c62..2e21a31 100644 --- a/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnSpeechDetector.java +++ b/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnSpeechDetector.java @@ -71,6 +71,11 @@ public final class RnnSpeechDetector implements SpeechProbabilityDetector { } } + /** See {@link RnnVad#resetNetwork()}. */ + public void resetNetwork() { + vad.resetNetwork(); + } + /** Appends samples (taking every {@code step}-th) and runs the detector per full frame. */ private void feed(float[] src, int offset, int length, int step) { for (int i = offset; i < offset + length; i += step) { diff --git a/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnVad.java b/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnVad.java index 69b54b3..5d649d9 100644 --- a/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnVad.java +++ b/ts3-client/core/src/main/java/com/ts3client/audio/vad/RnnVad.java @@ -59,6 +59,14 @@ public final class RnnVad { probability = 0.0f; } + /** + * Clears only the network's recurrent state, keeping the feature history. AGC2 does this + * every 1.5 s so its detector cannot latch. + */ + public void resetNetwork() { + network.reset(); + } + /** Speech probability in [0, 1] from the most recent {@link #process} call. */ public float probability() { return probability; diff --git a/ts3-client/core/src/test/java/com/ts3client/audio/processing/AudioProcessorTest.java b/ts3-client/core/src/test/java/com/ts3client/audio/processing/AudioProcessorTest.java new file mode 100644 index 0000000..eb45b17 --- /dev/null +++ b/ts3-client/core/src/test/java/com/ts3client/audio/processing/AudioProcessorTest.java @@ -0,0 +1,174 @@ +package com.ts3client.audio.processing; + +import com.ts3client.audio.processing.ns.SuppressionLevel; +import org.junit.jupiter.api.Test; + +import java.io.IOException; +import java.io.InputStream; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.util.Random; +import java.util.Set; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +/** + * Checks the chain against WebRTC itself. The {@code golden-*.s16} resources are + * {@link #testSignal()} run through upstream APM (webrtc-audio-processing 2.1 for noise + * suppression, 1.3 for the transient suppressor, which 2.x no longer ships), configured as + * TS3 configures it, and stored as 16-bit PCM. + */ +class AudioProcessorTest { + + private static final int RATE = AudioProcessor.SAMPLE_RATE; + private static final int FRAME = AudioProcessor.FRAME_SIZE; + + /** 10 ms frames in which the test signal has a key click, reported as key presses. */ + static final int[] KEY_FRAMES = {66, 72, 79, 85, 90, 97}; + private static final Set KEYS = Set.of(66, 72, 79, 85, 90, 97); + + /** + * One second of microphone-like input: quiet pink noise throughout, a voiced harmonic + * burst from 0.3 to 0.6 s, and key clicks after it. + */ + static float[] testSignal() { + Random random = new Random(11); + float[] x = new float[RATE]; + double b0 = 0, b1 = 0, b2 = 0; + for (int i = 0; i < x.length; i++) { + double w = random.nextGaussian(); + b0 = 0.99765 * b0 + w * 0.0990460; + b1 = 0.96300 * b1 + w * 0.2965164; + b2 = 0.57000 * b2 + w * 1.0526913; + double noise = (b0 + b1 + b2 + w * 0.1848) * 0.0012; + + double t = (double) i / RATE; + double voice = 0; + if (t >= 0.3 && t < 0.6) { + double envelope = Math.sin(Math.PI * (t - 0.3) / 0.3); + for (int h = 1; h <= 20; h++) { + voice += Math.sin(2 * Math.PI * 140 * h * t + h) / h; + } + voice *= 0.08 * envelope; + } + x[i] = (float) (noise + voice); + } + for (int frame : KEY_FRAMES) { + int start = frame * FRAME + 100; + for (int k = 0; k < 480; k++) { + x[start + k] += (float) (random.nextGaussian() * 0.1 * Math.exp(-k / 96.0)); + } + } + return x; + } + + private static float[] run(AudioProcessor processor, float[] input, Set keyFrames) { + float[] x = input.clone(); + float[] frame = new float[FRAME]; + for (int f = 0; f * FRAME + FRAME <= x.length; f++) { + if (keyFrames.contains(f)) processor.keyPressed(); + System.arraycopy(x, f * FRAME, frame, 0, FRAME); + processor.process(frame, FRAME); + System.arraycopy(frame, 0, x, f * FRAME, FRAME); + } + return x; + } + + private static void assertMatchesGolden(String name, float[] actual) throws IOException { + short[] golden = loadGolden(name); + assertEquals(golden.length, actual.length); + double maxDiff = 0; + double sumSquares = 0; + for (int i = 0; i < actual.length; i++) { + double d = actual[i] * 32768.0 - golden[i]; + maxDiff = Math.max(maxDiff, Math.abs(d)); + sumSquares += d * d; + } + double rms = Math.sqrt(sumSquares / actual.length); + // The reference is quantised to 16 bits, and float rounding differs slightly from + // upstream's: anything within a few LSB is the same signal. + assertTrue(maxDiff <= 4, name + ": max deviation " + maxDiff + " LSB"); + assertTrue(rms <= 0.6, name + ": RMS deviation " + rms + " LSB"); + } + + private static short[] loadGolden(String name) throws IOException { + try (InputStream in = AudioProcessorTest.class.getResourceAsStream(name)) { + assertNotNull(in, name); + ByteBuffer bytes = ByteBuffer.wrap(in.readAllBytes()).order(ByteOrder.LITTLE_ENDIAN); + short[] out = new short[bytes.remaining() / 2]; + bytes.asShortBuffer().get(out); + return out; + } + } + + @Test + void everythingOffPassesAudioThroughUntouched() { + float[] x = testSignal(); + assertArrayEquals(x, run(new AudioProcessor(), x, KEYS)); + } + + @Test + void noiseSuppressionMatchesWebRtcAtEveryLevel() throws IOException { + for (int level = 0; level <= 3; level++) { + AudioProcessor p = new AudioProcessor(); + p.setNoiseSuppression(true); + p.setSuppressionLevel(SuppressionLevel.fromDenoiserLevel(level)); + assertMatchesGolden("golden-ns" + level + ".s16", run(p, testSignal(), Set.of())); + } + } + + @Test + void typingAttenuationMatchesWebRtc() throws IOException { + AudioProcessor p = new AudioProcessor(); + p.setNoiseSuppression(true); + p.setSuppressionLevel(SuppressionLevel.DB_12); + p.setTransientSuppression(true); + assertMatchesGolden("golden-ns1-typing.s16", run(p, testSignal(), KEYS)); + } + + @Test + void typingAttenuationOnlyDelaysWithoutKeyPresses() { + float[] x = testSignal(); + AudioProcessor p = new AudioProcessor(); + p.setTransientSuppression(true); + float[] y = run(p, x, Set.of()); + int delay = 1024 - FRAME; + for (int i = delay; i < x.length; i++) { + assertEquals(x[i - delay], y[i], 0.0f); + } + } + + @Test + void gainControlRaisesQuietSpeechAndLeavesSilenceAlone() { + AudioProcessor p = new AudioProcessor(); + p.setGainControl(true); + float[] silence = run(p, new float[RATE], Set.of()); + for (float s : silence) assertEquals(0.0f, s, 0.0f); + + // 3 s of quiet voiced "words" at about -40 dBFS over a -70 dBFS noise floor. The + // pauses matter: without them AGC2 would take the voice for the noise floor. + Random random = new Random(5); + float[] quiet = new float[RATE * 3]; + for (int i = 0; i < quiet.length; i++) { + double t = (double) i / RATE; + double v = 0; + if (t % 0.5 < 0.3) { + for (int h = 1; h <= 20; h++) v += Math.sin(2 * Math.PI * 130 * h * t + h) / h; + v *= 0.004 * Math.sin(Math.PI * (t % 0.5) / 0.3); + } + quiet[i] = (float) (v + random.nextGaussian() * 0.0003); + } + float[] out = run(p, quiet, Set.of()); + assertTrue(rms(out, RATE * 2, RATE * 3) > 4 * rms(quiet, RATE * 2, RATE * 3), + "quiet speech should be raised by more than 12 dB"); + } + + private static double rms(float[] x, int from, int to) { + double sum = 0; + for (int i = from; i < to; i++) sum += x[i] * x[i]; + return Math.sqrt(sum / (to - from)); + } +} diff --git a/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns0.s16 b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns0.s16 new file mode 100644 index 0000000..edf3a4d Binary files /dev/null and b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns0.s16 differ diff --git a/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns1-typing.s16 b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns1-typing.s16 new file mode 100644 index 0000000..53f00f8 Binary files /dev/null and b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns1-typing.s16 differ diff --git a/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns1.s16 b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns1.s16 new file mode 100644 index 0000000..23bf450 Binary files /dev/null and b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns1.s16 differ diff --git a/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns2.s16 b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns2.s16 new file mode 100644 index 0000000..b2f2266 Binary files /dev/null and b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns2.s16 differ diff --git a/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns3.s16 b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns3.s16 new file mode 100644 index 0000000..c586d0a Binary files /dev/null and b/ts3-client/core/src/test/resources/com/ts3client/audio/processing/golden-ns3.s16 differ