package com.sherpaonnxofflinestt.managers import com.k2fsa.sherpa.onnx.OfflineSpeechDenoiser import com.k2fsa.sherpa.onnx.OfflineSpeechDenoiserConfig import com.k2fsa.sherpa.onnx.OfflineSpeechDenoiserGtcrnModelConfig import com.k2fsa.sherpa.onnx.OfflineSpeechDenoiserModelConfig import com.sherpaonnxofflinestt.STTConstants import com.sherpaonnxofflinestt.utils.STTLogger import java.io.File /** * Manages GTCRN speech denoiser initialization and processing */ class DenoiserManager { private var denoiser: OfflineSpeechDenoiser? = null var isEnabled: Boolean = false private set /** * Initialize the denoiser with model path * @param modelPath Path to GTCRN model file * @return true if initialization successful */ fun initialize(modelPath: String): Boolean { if (modelPath.isEmpty() || !File(modelPath).exists()) { STTLogger.Denoiser.w("Model path invalid or not found: $modelPath") return false } return try { val gtcrnConfig = OfflineSpeechDenoiserGtcrnModelConfig(model = modelPath) val modelConfig = OfflineSpeechDenoiserModelConfig( gtcrn = gtcrnConfig, numThreads = STTConstants.DEFAULT_AUXILIARY_NUM_THREADS, debug = false, provider = STTConstants.DEFAULT_PROVIDER ) val config = OfflineSpeechDenoiserConfig(model = modelConfig) denoiser = OfflineSpeechDenoiser(assetManager = null, config = config) isEnabled = true STTLogger.Denoiser.i("Denoiser initialized successfully") true } catch (e: Exception) { STTLogger.Denoiser.e("Failed to initialize denoiser: ${e.message}", e) isEnabled = false false } } /** * Enable or disable the denoiser * @param enabled Whether to enable denoising * @return The new enabled state, or null if denoiser not initialized */ fun setEnabled(enabled: Boolean): Boolean? { if (denoiser == null) { STTLogger.Denoiser.w("Cannot set enabled: denoiser not initialized") return null } isEnabled = enabled STTLogger.Denoiser.i("Denoiser ${if (enabled) "enabled" else "disabled"}") return isEnabled } /** * Maximum chunk size for denoising. * Longer audio is processed in chunks to prevent model state drift. */ private val maxChunkSamples = STTConstants.MAX_DENOISER_CHUNK_SAMPLES /** * Process audio samples through the denoiser * For long audio (>5s), processes in chunks to prevent state drift * @param samples Input audio samples (float array) * @param sampleRate Sample rate in Hz * @return Denoised audio samples, or original samples if denoising fails */ fun process(samples: FloatArray, sampleRate: Int): FloatArray { if (!isEnabled || denoiser == null) { return samples } return try { // For short audio, process directly if (samples.size <= maxChunkSamples) { val denoisedAudio = denoiser!!.run(samples, sampleRate) return denoisedAudio.samples } // For long audio, process in chunks to prevent state drift STTLogger.Denoiser.i("Processing ${samples.size} samples in chunks (${samples.size / maxChunkSamples + 1} chunks)") val result = FloatArray(samples.size) var offset = 0 var chunkIndex = 0 while (offset < samples.size) { val chunkEnd = minOf(offset + maxChunkSamples, samples.size) val chunkSize = chunkEnd - offset // Extract chunk val chunk = samples.copyOfRange(offset, chunkEnd) // Process chunk val denoisedChunk = denoiser!!.run(chunk, sampleRate) // Copy to result (handle potential size mismatch) val copySize = minOf(denoisedChunk.samples.size, chunkSize) denoisedChunk.samples.copyInto(result, offset, 0, copySize) chunkIndex++ offset = chunkEnd } STTLogger.Denoiser.i("Processed $chunkIndex chunks successfully") result } catch (e: Exception) { STTLogger.Denoiser.w("Denoising failed: ${e.message}") samples } } /** * Check if denoiser is initialized */ fun isInitialized(): Boolean = denoiser != null /** * Release denoiser resources */ fun release() { denoiser?.release() denoiser = null isEnabled = false STTLogger.Denoiser.i("Denoiser released") } }