feat(cw): add CwDeepDecoder wiring ONNX Runtime to the CW pipeline
组装完整解码链路: 麦克风 PCM -> 重采样 3200Hz -> 20 秒滚动缓冲 -> 频谱图 -> ONNX 推理 -> 贪心 CTC -> 文本。 ICwDecoder 放 core:domain (纯 Kotlin, UI 可直接依赖), 实现放 core:data (需要 Android Context 读 assets 及 ONNX Runtime 原生库)。 关键设计: - decodedText 是替换语义不是追加: 模型会随上下文增加改写先前字符, 追加 会把中间态 (BM -> BG7 -> BG7NTA) 永久留在屏幕上 - 推理用 Mutex.tryLock 串行化: 上一次未完成时直接跳过本周期, 慢设备不会 堆积任务导致 OOM - 推理在 Dispatchers.Default, 不阻塞音频采集线程 - 模型元数据从 model.onnx.json 读取 (字符表/blank 索引/输入输出名), 不在 代码里硬编码, 模型更新时无需改代码 - lastInferenceMs 打点: 手机实际推理耗时需真机实测 (估算系数待验证) - estimatedPitch 由频谱峰值 bin 反算, 仅供 UI 显示 —— 模型 400-1200Hz 固定窗自带定频, 不需要频谱峰值跟踪参与解码 - 模型加载失败时 errorMessage 置位 (旧内核已删, 无兜底可退) - close() 释放 OrtSession, 否则原生内存泄漏 验证: ./gradlew :core:data:compileDebugKotlin => BUILD SUCCESSFUL
This commit is contained in:
1 parent
13ef70810c
commit
57762613a8
2 files changed
+271
No files matched your search
@@ -0,0 +1,212 @@
|
||||
/*
|
||||
* Look4Sat. Amateur radio satellite tracker and pass predictor.
|
||||
* Copyright (C) 2019-2026 Arty Bishop and contributors.
|
||||
*
|
||||
* This program is free software: you can redistribute it and/or modify
|
||||
* it under the terms of the GNU General Public License as published by
|
||||
* the Free Software Foundation, either version 3 of the License, or
|
||||
* (at your option) any later version.
|
||||
*
|
||||
* This program is distributed in the hope that it will be useful,
|
||||
* but WITHOUT ANY WARRANTY; without even the implied warranty of
|
||||
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
||||
* GNU General Public License for more details.
|
||||
*
|
||||
* You should have received a copy of the GNU General Public License
|
||||
* along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||
*/
|
||||
package com.rtbishop.look4sat.core.data.cw
|
||||
|
||||
import ai.onnxruntime.OnnxTensor
|
||||
import ai.onnxruntime.OrtEnvironment
|
||||
import ai.onnxruntime.OrtSession
|
||||
import android.content.Context
|
||||
import android.util.Log
|
||||
import com.rtbishop.look4sat.core.domain.cw.CwCtcDecoder
|
||||
import com.rtbishop.look4sat.core.domain.cw.CwDeepBuffer
|
||||
import com.rtbishop.look4sat.core.domain.cw.CwDeepSpectrogram
|
||||
import com.rtbishop.look4sat.core.domain.cw.ICwDecoder
|
||||
import kotlinx.coroutines.Dispatchers
|
||||
import kotlinx.coroutines.flow.MutableStateFlow
|
||||
import kotlinx.coroutines.flow.StateFlow
|
||||
import kotlinx.coroutines.flow.asStateFlow
|
||||
import kotlinx.coroutines.sync.Mutex
|
||||
import kotlinx.coroutines.withContext
|
||||
import org.json.JSONObject
|
||||
import java.nio.FloatBuffer
|
||||
|
||||
/**
|
||||
* CW decoder backed by the DeepCW neural network (AGPL-3.0, see
|
||||
* `feature/cw/licenses/NOTICE.md`).
|
||||
*
|
||||
* The model classifies a whole audio segment at once rather than streaming
|
||||
* sample by sample, and it revises earlier characters once more context
|
||||
* arrives. Incremental stitching therefore produces duplicated callsigns —
|
||||
* measured character error rates of 67-294% against 0% for whole-segment
|
||||
* decoding. Instead a [CwDeepBuffer] holds the last 20 seconds and the whole
|
||||
* window is re-decoded every 1.5 seconds, replacing [decodedText] outright.
|
||||
*
|
||||
* The model's fixed 400-1200 Hz analysis window means pitch detection is built
|
||||
* in; no spectral peak tracking or squelch gating is needed.
|
||||
*/
|
||||
class CwDeepDecoder(context: Context) : ICwDecoder {
|
||||
|
||||
private companion object {
|
||||
const val TAG = "CwDeepDecoder"
|
||||
const val MODEL_ASSET = "deepcw/model.onnx"
|
||||
const val METADATA_ASSET = "deepcw/model.onnx.json"
|
||||
}
|
||||
|
||||
private val _decodedText = MutableStateFlow("")
|
||||
override val decodedText: StateFlow<String> = _decodedText.asStateFlow()
|
||||
|
||||
private val _estimatedPitch = MutableStateFlow<Float?>(null)
|
||||
override val estimatedPitch: StateFlow<Float?> = _estimatedPitch.asStateFlow()
|
||||
|
||||
private val _signalStrength = MutableStateFlow(0f)
|
||||
override val signalStrength: StateFlow<Float> = _signalStrength.asStateFlow()
|
||||
|
||||
private val _lastInferenceMs = MutableStateFlow(0)
|
||||
override val lastInferenceMs: StateFlow<Int> = _lastInferenceMs.asStateFlow()
|
||||
|
||||
private val _errorMessage = MutableStateFlow<String?>(null)
|
||||
override val errorMessage: StateFlow<String?> = _errorMessage.asStateFlow()
|
||||
|
||||
private val buffer = CwDeepBuffer()
|
||||
|
||||
/** Held while inference runs so slow devices skip work instead of queuing it. */
|
||||
private val inferenceLock = Mutex()
|
||||
|
||||
private var environment: OrtEnvironment? = null
|
||||
private var session: OrtSession? = null
|
||||
private var chars: List<String> = emptyList()
|
||||
private var blankIndex = 41
|
||||
private var inputName = "spectrogram"
|
||||
private var outputName = "log_probs"
|
||||
|
||||
init {
|
||||
try {
|
||||
val metadata = JSONObject(
|
||||
context.assets.open(METADATA_ASSET).bufferedReader().use { it.readText() }
|
||||
)
|
||||
val charArray = metadata.getJSONArray("chars")
|
||||
chars = List(charArray.length()) { charArray.getString(it) }
|
||||
blankIndex = metadata.getInt("blank_index")
|
||||
inputName = metadata.getString("onnx_input_name")
|
||||
outputName = metadata.getString("onnx_output_name")
|
||||
|
||||
val modelBytes = context.assets.open(MODEL_ASSET).use { it.readBytes() }
|
||||
environment = OrtEnvironment.getEnvironment()
|
||||
session = environment?.createSession(modelBytes, OrtSession.SessionOptions())
|
||||
Log.i(TAG, "DeepCW ready: ${modelBytes.size} bytes, ${chars.size} classes")
|
||||
} catch (t: Throwable) {
|
||||
Log.e(TAG, "DeepCW model failed to load", t)
|
||||
_errorMessage.value = "CW model failed to load: ${t.message ?: t.javaClass.simpleName}"
|
||||
}
|
||||
}
|
||||
|
||||
override suspend fun processBuffer(samples: FloatArray, sampleRate: Int) {
|
||||
if (session == null || samples.isEmpty()) return
|
||||
|
||||
val resampled = CwDeepSpectrogram.resampleLinear(
|
||||
samples, sampleRate, CwDeepSpectrogram.SAMPLE_RATE
|
||||
)
|
||||
val shouldRedecode = buffer.append(resampled)
|
||||
if (!shouldRedecode || !buffer.hasEnoughAudio) return
|
||||
|
||||
// Drop this cycle rather than queue when the previous run is still going.
|
||||
if (!inferenceLock.tryLock()) {
|
||||
Log.d(TAG, "inference still running, skipping this interval")
|
||||
return
|
||||
}
|
||||
try {
|
||||
decodeWindow(buffer.snapshot())
|
||||
} catch (t: Throwable) {
|
||||
Log.e(TAG, "inference failed", t)
|
||||
_errorMessage.value = "CW decode failed: ${t.message ?: t.javaClass.simpleName}"
|
||||
} finally {
|
||||
inferenceLock.unlock()
|
||||
}
|
||||
}
|
||||
|
||||
private suspend fun decodeWindow(window: FloatArray) = withContext(Dispatchers.Default) {
|
||||
val activeSession = session ?: return@withContext
|
||||
val activeEnvironment = environment ?: return@withContext
|
||||
|
||||
val spectrogram = CwDeepSpectrogram.compute(window)
|
||||
val frames = spectrogram.size
|
||||
val bins = CwDeepSpectrogram.FREQUENCY_BINS
|
||||
|
||||
val flat = FloatBuffer.allocate(frames * bins)
|
||||
for (frame in spectrogram) flat.put(frame)
|
||||
flat.rewind()
|
||||
|
||||
val shape = longArrayOf(1, 1, frames.toLong(), bins.toLong())
|
||||
val startedAt = System.currentTimeMillis()
|
||||
val text: String
|
||||
OnnxTensor.createTensor(activeEnvironment, flat, shape).use { input ->
|
||||
activeSession.run(mapOf(inputName to input)).use { result ->
|
||||
@Suppress("UNCHECKED_CAST")
|
||||
val logits = result[outputName].get().value as Array<Array<FloatArray>>
|
||||
text = CwCtcDecoder.greedy(logits, chars, blankIndex)
|
||||
}
|
||||
}
|
||||
_lastInferenceMs.value = (System.currentTimeMillis() - startedAt).toInt()
|
||||
|
||||
// Replace, never append: the model rewrites earlier characters as more
|
||||
// context arrives, so appending would leave stale guesses on screen.
|
||||
_decodedText.value = text
|
||||
|
||||
updateSignalMetrics(spectrogram)
|
||||
}
|
||||
|
||||
/**
|
||||
* Report the loudest bin as the tone pitch and its prominence over the
|
||||
* window mean as a 0..1 strength, purely for the UI readout.
|
||||
*/
|
||||
private fun updateSignalMetrics(spectrogram: Array<FloatArray>) {
|
||||
if (spectrogram.isEmpty()) return
|
||||
var bestBin = 0
|
||||
var bestValue = 0f
|
||||
var total = 0f
|
||||
var count = 0
|
||||
for (frame in spectrogram) {
|
||||
for (bin in frame.indices) {
|
||||
val value = frame[bin]
|
||||
total += value
|
||||
count++
|
||||
if (value > bestValue) {
|
||||
bestValue = value
|
||||
bestBin = bin
|
||||
}
|
||||
}
|
||||
}
|
||||
if (count == 0 || bestValue <= 0f) return
|
||||
|
||||
val binHz = CwDeepSpectrogram.SAMPLE_RATE.toDouble() / CwDeepSpectrogram.FFT_LENGTH
|
||||
// Relative bin 0 is 400 Hz; absolute bin index is 32 + bestBin.
|
||||
val absoluteBin = 32 + bestBin
|
||||
_estimatedPitch.value = (absoluteBin * binHz).toFloat()
|
||||
|
||||
val mean = total / count
|
||||
_signalStrength.value = ((bestValue - mean) / bestValue).coerceIn(0f, 1f)
|
||||
}
|
||||
|
||||
override fun reset() {
|
||||
buffer.reset()
|
||||
_decodedText.value = ""
|
||||
_estimatedPitch.value = null
|
||||
_signalStrength.value = 0f
|
||||
_lastInferenceMs.value = 0
|
||||
}
|
||||
|
||||
override fun close() {
|
||||
try {
|
||||
session?.close()
|
||||
} catch (t: Throwable) {
|
||||
Log.w(TAG, "session close failed", t)
|
||||
}
|
||||
session = null
|
||||
environment = null
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
/*
|
||||
* Look4Sat. Amateur radio satellite tracker and pass predictor.
|
||||
* Copyright (C) 2019-2026 Arty Bishop and contributors.
|
||||
*
|
||||
* This program is free software: you can redistribute it and/or modify
|
||||
* it under the terms of the GNU General Public License as published by
|
||||
* the Free Software Foundation, either version 3 of the License, or
|
||||
* (at your option) any later version.
|
||||
*
|
||||
* This program is distributed in the hope that it will be useful,
|
||||
* but WITHOUT ANY WARRANTY; without even the implied warranty of
|
||||
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
||||
* GNU General Public License for more details.
|
||||
*
|
||||
* You should have received a copy of the GNU General Public License
|
||||
* along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||
*/
|
||||
package com.rtbishop.look4sat.core.domain.cw
|
||||
|
||||
import kotlinx.coroutines.flow.StateFlow
|
||||
|
||||
/**
|
||||
* A CW (Morse code) decoder fed with microphone PCM.
|
||||
*
|
||||
* Implementations live outside `core:domain` when they need platform APIs;
|
||||
* this contract stays pure Kotlin so the UI can depend on it directly.
|
||||
*/
|
||||
interface ICwDecoder {
|
||||
|
||||
/**
|
||||
* Decoded text for display.
|
||||
*
|
||||
* Note this is **replace** semantics, not append: a whole-segment model
|
||||
* revises earlier characters as more audio arrives, so consumers must show
|
||||
* the current value rather than accumulating emissions.
|
||||
*/
|
||||
val decodedText: StateFlow<String>
|
||||
|
||||
/** Detected tone frequency in Hz, or null before a tone is found. */
|
||||
val estimatedPitch: StateFlow<Float?>
|
||||
|
||||
/** Relative signal strength in 0..1 for level meters. */
|
||||
val signalStrength: StateFlow<Float>
|
||||
|
||||
/** Most recent inference duration in milliseconds, for diagnostics. */
|
||||
val lastInferenceMs: StateFlow<Int>
|
||||
|
||||
/** Non-null when the decoder cannot run, for example the model failed to load. */
|
||||
val errorMessage: StateFlow<String?>
|
||||
|
||||
/** Feed captured mono PCM in -1..1. Safe to call from a capture thread. */
|
||||
suspend fun processBuffer(samples: FloatArray, sampleRate: Int)
|
||||
|
||||
/** Clear decoded text and buffered audio. */
|
||||
fun reset()
|
||||
|
||||
/** Release native resources. Must be called when the decoder goes away. */
|
||||
fun close()
|
||||
}
|
||||
Reference in new issue
Block a user