feat(cw): add pure-Kotlin DeepCW spectrogram front-end
DeepCW 音频前处理的 Kotlin 实现, 逐步对齐上游 Python 参考实现 (deepcw-engine examples/python/decode_morse.py): 线性重采样 -> 3200Hz, 反射 padding(128), periodic Hann(256), radix-2 FFT, 取 bin [32,97) 共 65 个 (400-1200Hz), log1p 归一化 实现要点: - Hann 窗用 periodic 变体 (numpy np.hanning(N+1)[:-1]), 不是 symmetric, 否则频谱与参考实现有偏差 - 反射 padding 不重复边界样本 (numpy mode="reflect" 语义) - 自带 radix-2 Cooley-Tukey FFT, 不引入第三方 DSP 依赖 (FFT_LENGTH 为 2 的幂, 无需补零分支) - 放 core:domain: 无 Android 依赖, 可 JVM 单测, 保持 KMP 可迁移 模型 400-1200Hz 固定频窗意味着定频内建, 无需频谱峰值跟踪与 squelch 门控 —— 这正是旧内核最大的两个坑。 验证: ./gradlew :core:domain:test --tests '*CwDeepSpectrogramTest*' => tests="8" skipped="0" failures="0" errors="0" 含断言: 700Hz 正弦峰值落在相对 bin 24; 8000->3200Hz 重采样后频率不变
This commit is contained in:
1 parent
b8113fc757
commit
7c22bf6ac1
2 files changed
+318
No files matched your search
@@ -0,0 +1,208 @@
|
||||
/*
|
||||
* 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 kotlin.math.PI
|
||||
import kotlin.math.ceil
|
||||
import kotlin.math.cos
|
||||
import kotlin.math.floor
|
||||
import kotlin.math.ln1p
|
||||
import kotlin.math.roundToInt
|
||||
import kotlin.math.sqrt
|
||||
|
||||
/**
|
||||
* Audio front-end for the DeepCW model: turns PCM samples into the
|
||||
* `[time, frequency]` log-magnitude spectrogram the network expects.
|
||||
*
|
||||
* Mirrors the upstream Python reference (deepcw-engine
|
||||
* `examples/python/decode_morse.py`) step for step:
|
||||
*
|
||||
* resample -> 3200 Hz, reflect-pad by fft/2, periodic Hann window of 256,
|
||||
* real FFT, keep bins [32, 97) i.e. 400-1200 Hz, then log1p.
|
||||
*
|
||||
* The model's fixed 400-1200 Hz window means pitch detection is built in —
|
||||
* no spectral peak tracking or squelch gating is needed on our side.
|
||||
*/
|
||||
object CwDeepSpectrogram {
|
||||
|
||||
/** Model input sample rate, from `model.onnx.json`. */
|
||||
const val SAMPLE_RATE = 3200
|
||||
|
||||
/** FFT window length in samples. */
|
||||
const val FFT_LENGTH = 256
|
||||
|
||||
/** Hop between consecutive frames; 48/3200 = 15.0 ms per frame. */
|
||||
const val HOP_LENGTH = 48
|
||||
|
||||
private const val MIN_FREQ_HZ = 400.0
|
||||
private const val MAX_FREQ_HZ = 1200.0
|
||||
|
||||
/** Number of frequency bins the model expects. */
|
||||
const val FREQUENCY_BINS = 65
|
||||
|
||||
/** Milliseconds of audio represented by one output frame. */
|
||||
const val MS_PER_FRAME = 1000.0 * HOP_LENGTH / SAMPLE_RATE
|
||||
|
||||
private val hannWindow: FloatArray = FloatArray(FFT_LENGTH) { i ->
|
||||
// numpy: np.hanning(N + 1)[:-1] — the periodic (not symmetric) variant.
|
||||
(0.5 - 0.5 * cos(2.0 * PI * i / FFT_LENGTH)).toFloat()
|
||||
}
|
||||
|
||||
/**
|
||||
* Inclusive-exclusive bin range covering [minHz, maxHz].
|
||||
* Returns `start to stop`, matching the reference's `frequency_bin_range`.
|
||||
*/
|
||||
fun frequencyBinRange(
|
||||
sampleRate: Int,
|
||||
fftLength: Int,
|
||||
minHz: Double,
|
||||
maxHz: Double
|
||||
): Pair<Int, Int> {
|
||||
val binHz = sampleRate.toDouble() / fftLength
|
||||
val start = ceil(minHz / binHz).toInt()
|
||||
val stop = floor(maxHz / binHz).toInt() + 1
|
||||
return start to stop
|
||||
}
|
||||
|
||||
/**
|
||||
* Linear-interpolation resampler. Deliberately dependency-light and
|
||||
* identical to the reference implementation so spectrograms match.
|
||||
*/
|
||||
fun resampleLinear(audio: FloatArray, sourceRate: Int, targetRate: Int): FloatArray {
|
||||
if (sourceRate == targetRate || audio.isEmpty()) return audio
|
||||
val targetLength = (audio.size.toDouble() * targetRate / sourceRate).roundToInt()
|
||||
val out = FloatArray(targetLength)
|
||||
val ratio = sourceRate.toDouble() / targetRate
|
||||
for (i in 0 until targetLength) {
|
||||
val position = i * ratio
|
||||
val left = floor(position).toInt()
|
||||
val right = minOf(left + 1, audio.size - 1)
|
||||
val fraction = (position - left).toFloat()
|
||||
out[i] = audio[left] * (1f - fraction) + audio[right] * fraction
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
/**
|
||||
* Build the log-magnitude spectrogram. Input must already be at
|
||||
* [SAMPLE_RATE]; use [resampleLinear] first when it is not.
|
||||
*
|
||||
* @return `[frames][FREQUENCY_BINS]` values, all non-negative.
|
||||
*/
|
||||
fun compute(audio: FloatArray): Array<FloatArray> {
|
||||
require(audio.size >= FFT_LENGTH) {
|
||||
"audio is too short for fftLength=$FFT_LENGTH, got ${audio.size}"
|
||||
}
|
||||
|
||||
val (startBin, stopBin) = frequencyBinRange(
|
||||
SAMPLE_RATE, FFT_LENGTH, MIN_FREQ_HZ, MAX_FREQ_HZ
|
||||
)
|
||||
val bins = stopBin - startBin
|
||||
require(bins == FREQUENCY_BINS) {
|
||||
"expected $FREQUENCY_BINS bins, computed $bins"
|
||||
}
|
||||
|
||||
val padded = reflectPad(audio, FFT_LENGTH / 2)
|
||||
val frames = 1 + (padded.size - FFT_LENGTH) / HOP_LENGTH
|
||||
val result = Array(frames) { FloatArray(bins) }
|
||||
|
||||
val real = FloatArray(FFT_LENGTH)
|
||||
val imag = FloatArray(FFT_LENGTH)
|
||||
for (frame in 0 until frames) {
|
||||
val offset = frame * HOP_LENGTH
|
||||
for (i in 0 until FFT_LENGTH) {
|
||||
real[i] = padded[offset + i] * hannWindow[i]
|
||||
imag[i] = 0f
|
||||
}
|
||||
fftInPlace(real, imag)
|
||||
val row = result[frame]
|
||||
for (bin in startBin until stopBin) {
|
||||
val magnitude = sqrt(real[bin] * real[bin] + imag[bin] * imag[bin])
|
||||
row[bin - startBin] = ln1p(magnitude.toDouble()).toFloat()
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
/**
|
||||
* numpy `mode="reflect"`: mirrors around the edge samples without
|
||||
* repeating them, so [1,2,3] padded by 2 becomes [3,2,1,2,3,2,1].
|
||||
*/
|
||||
private fun reflectPad(audio: FloatArray, pad: Int): FloatArray {
|
||||
if (pad == 0) return audio
|
||||
val out = FloatArray(audio.size + 2 * pad)
|
||||
for (i in 0 until pad) out[i] = audio[pad - i]
|
||||
audio.copyInto(out, pad)
|
||||
val last = audio.size - 1
|
||||
for (i in 0 until pad) out[pad + audio.size + i] = audio[last - 1 - i]
|
||||
return out
|
||||
}
|
||||
|
||||
/**
|
||||
* Iterative radix-2 Cooley-Tukey FFT. [FFT_LENGTH] is a power of two, so
|
||||
* no padding case is needed. Only the first half of the output is read by
|
||||
* [compute], which is the real-input equivalent of numpy's `rfft`.
|
||||
*/
|
||||
private fun fftInPlace(real: FloatArray, imag: FloatArray) {
|
||||
val n = real.size
|
||||
|
||||
// Bit-reversal permutation.
|
||||
var j = 0
|
||||
for (i in 1 until n) {
|
||||
var bit = n shr 1
|
||||
while (j and bit != 0) {
|
||||
j = j xor bit
|
||||
bit = bit shr 1
|
||||
}
|
||||
j = j or bit
|
||||
if (i < j) {
|
||||
var tmp = real[i]; real[i] = real[j]; real[j] = tmp
|
||||
tmp = imag[i]; imag[i] = imag[j]; imag[j] = tmp
|
||||
}
|
||||
}
|
||||
|
||||
var length = 2
|
||||
while (length <= n) {
|
||||
val angle = -2.0 * PI / length
|
||||
val wReal = cos(angle).toFloat()
|
||||
val wImag = kotlin.math.sin(angle).toFloat()
|
||||
var i = 0
|
||||
while (i < n) {
|
||||
var curReal = 1f
|
||||
var curImag = 0f
|
||||
for (k in 0 until length / 2) {
|
||||
val evenReal = real[i + k]
|
||||
val evenImag = imag[i + k]
|
||||
val oddReal = real[i + k + length / 2]
|
||||
val oddImag = imag[i + k + length / 2]
|
||||
val mulReal = oddReal * curReal - oddImag * curImag
|
||||
val mulImag = oddReal * curImag + oddImag * curReal
|
||||
real[i + k] = evenReal + mulReal
|
||||
imag[i + k] = evenImag + mulImag
|
||||
real[i + k + length / 2] = evenReal - mulReal
|
||||
imag[i + k + length / 2] = evenImag - mulImag
|
||||
val nextReal = curReal * wReal - curImag * wImag
|
||||
curImag = curReal * wImag + curImag * wReal
|
||||
curReal = nextReal
|
||||
}
|
||||
i += length
|
||||
}
|
||||
length = length shl 1
|
||||
}
|
||||
}
|
||||
}
|
||||
+110
@@ -0,0 +1,110 @@
|
||||
/*
|
||||
* 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 org.junit.Assert.assertEquals
|
||||
import org.junit.Assert.assertTrue
|
||||
import org.junit.Test
|
||||
import kotlin.math.PI
|
||||
import kotlin.math.abs
|
||||
import kotlin.math.sin
|
||||
|
||||
/**
|
||||
* Verifies the DeepCW front-end against the upstream Python reference
|
||||
* implementation (deepcw-engine examples/python/decode_morse.py).
|
||||
*
|
||||
* Model metadata: sampleRate 3200, fftLength 256, hopLength 48,
|
||||
* 400-1200 Hz -> 65 bins, log1p normalization.
|
||||
*/
|
||||
class CwDeepSpectrogramTest {
|
||||
|
||||
@Test
|
||||
fun frequencyBinRange_matchesModelMetadata() {
|
||||
// binHz = 3200/256 = 12.5; start = ceil(400/12.5) = 32; stop = floor(1200/12.5)+1 = 97
|
||||
val (start, stop) = CwDeepSpectrogram.frequencyBinRange(3200, 256, 400.0, 1200.0)
|
||||
assertEquals(32, start)
|
||||
assertEquals(97, stop)
|
||||
assertEquals("metadata declares 65 frequency bins", 65, stop - start)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun compute_producesTimeBy65Matrix() {
|
||||
// 1 second at 3200 Hz. Reflect padding adds fft/2 on both sides,
|
||||
// so frames = 1 + (3200 + 256 - 256)/48 = 1 + 66 = 67
|
||||
val spec = CwDeepSpectrogram.compute(FloatArray(3200))
|
||||
assertEquals(67, spec.size)
|
||||
assertEquals(65, spec[0].size)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun compute_toneLandsInExpectedBin() {
|
||||
// 700 Hz -> absolute bin 700/12.5 = 56 -> relative index 56 - 32 = 24
|
||||
val audio = FloatArray(3200) { (0.6 * sin(2.0 * PI * 700.0 * it / 3200.0)).toFloat() }
|
||||
val spec = CwDeepSpectrogram.compute(audio)
|
||||
val middle = spec[spec.size / 2]
|
||||
val peak = middle.indices.maxByOrNull { middle[it] } ?: -1
|
||||
assertTrue("peak at index $peak, expected near 24", abs(peak - 24) <= 1)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun compute_appliesLog1pSoValuesAreNonNegative() {
|
||||
val audio = FloatArray(3200) { (0.6 * sin(2.0 * PI * 700.0 * it / 3200.0)).toFloat() }
|
||||
val spec = CwDeepSpectrogram.compute(audio)
|
||||
for (frame in spec) {
|
||||
for (v in frame) {
|
||||
assertTrue("log1p of a magnitude must be >= 0, got $v", v >= 0f)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
fun resampleLinear_convertsRateAndLength() {
|
||||
assertEquals(3200, CwDeepSpectrogram.resampleLinear(FloatArray(8000), 8000, 3200).size)
|
||||
assertEquals(3200, CwDeepSpectrogram.resampleLinear(FloatArray(44100), 44100, 3200).size)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun resampleLinear_sameRateIsIdentity() {
|
||||
val input = floatArrayOf(0.1f, 0.2f, 0.3f)
|
||||
val out = CwDeepSpectrogram.resampleLinear(input, 3200, 3200)
|
||||
assertEquals(3, out.size)
|
||||
assertEquals(0.2f, out[1], 1e-6f)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun resampleLinear_preservesToneFrequency() {
|
||||
// A 700 Hz tone sampled at 8000 Hz must still peak at bin 24 after
|
||||
// resampling to 3200 Hz — this is the path real microphone audio takes.
|
||||
val at8k = FloatArray(8000) { (0.6 * sin(2.0 * PI * 700.0 * it / 8000.0)).toFloat() }
|
||||
val at3200 = CwDeepSpectrogram.resampleLinear(at8k, 8000, 3200)
|
||||
val spec = CwDeepSpectrogram.compute(at3200)
|
||||
val middle = spec[spec.size / 2]
|
||||
val peak = middle.indices.maxByOrNull { middle[it] } ?: -1
|
||||
assertTrue("resampled tone peak at $peak, expected near 24", abs(peak - 24) <= 1)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun compute_rejectsAudioShorterThanFftLength() {
|
||||
try {
|
||||
CwDeepSpectrogram.compute(FloatArray(100))
|
||||
throw AssertionError("expected an exception for audio shorter than fftLength")
|
||||
} catch (expected: IllegalArgumentException) {
|
||||
// desired path
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user