feat(cw): add greedy CTC collapse for DeepCW output

把模型 log_probs [batch,time,42] 输出折叠成文本, 对齐参考实现
greedy_ctc_decode: 逐帧取 argmax, 去 blank, 再去连续重复标签。

blank 的作用: 两个相同字母之间的 blank 是保住真实双字母的关键
(例如 "5NN" 的两个 N), 否则会被折叠成一个。

验证:
./gradlew :core:domain:test --tests '*CwCtcDecoderTest*'
=> tests="7" skipped="0" failures="0" errors="0"
含断言: 连续重复折叠、blank 隔断保留双字母、呼号 "CQ BG7" 含空格与数字
This commit is contained in:
mckero committed 2026-08-12 12:07:50 +00:00
1 parent 7c22bf6ac1
commit 0d3e02c596
2 files changed
+151

No files matched your search

@@ -0,0 +1,61 @@
/*
* 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
/**
* Turns the model's `log_probs` output into text.
*
* Mirrors the reference `greedy_ctc_decode`: take the best class per frame,
* drop blanks, and collapse runs of the same label. A blank between two
* identical labels is what keeps a genuine double letter (for example the
* two N's in "5NN") from collapsing into one.
*/
object CwCtcDecoder {
/**
* @param logProbs `[batch, time, class]`; only batch 0 is read.
* @param chars class index to symbol, excluding the blank.
* @param blankIndex the CTC blank class (41 for this model).
*/
fun greedy(
logProbs: Array<Array<FloatArray>>,
chars: List<String>,
blankIndex: Int
): String {
if (logProbs.isEmpty()) return ""
val frames = logProbs[0]
val builder = StringBuilder()
var previous = -1
for (frame in frames) {
var best = 0
for (i in 1 until frame.size) {
if (frame[i] > frame[best]) best = i
}
if (best == blankIndex) {
previous = -1
continue
}
if (best != previous && best < chars.size) {
builder.append(chars[best])
}
previous = best
}
return builder.toString()
}
}
@@ -0,0 +1,90 @@
/*
* 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.Test
/**
* Greedy CTC collapse, matching the reference implementation's
* `greedy_ctc_decode`: drop blanks, then drop runs of the same label.
*
* Alphabet from `model.onnx.json` — 41 symbols plus blank at index 41.
*/
class CwCtcDecoderTest {
private val chars = listOf(
",", ".", "/", "0", "1", "2", "3", "4", "5", "6", "7", "8", "9", "?",
"A", "B", "C", "D", "E", "F", "G", "H", "I", "J", "K", "L", "M", "N",
"O", "P", "Q", "R", "S", "T", "U", "V", "W", "X", "Y", "Z", " "
)
private val blank = 41
/** Build a `[1, T, 42]` log-prob tensor whose argmax follows [path]. */
private fun logits(path: IntArray): Array<Array<FloatArray>> {
val frames = Array(path.size) { t ->
FloatArray(42) { -10f }.also { it[path[t]] = 0f }
}
return arrayOf(frames)
}
@Test
fun alphabetSizeMatchesModelMetadata() {
assertEquals("41 symbols + blank = 42 classes", 41, chars.size)
}
@Test
fun greedy_dropsRunsOfTheSameLabel() {
assertEquals("A", CwCtcDecoder.greedy(logits(intArrayOf(14, 14, 14)), chars, blank))
}
@Test
fun greedy_keepsRepeatsSeparatedByBlank() {
// A A <blank> A collapses to "AA": the blank breaks the run.
assertEquals("AA", CwCtcDecoder.greedy(logits(intArrayOf(14, 14, blank, 14)), chars, blank))
}
@Test
fun greedy_allBlanksYieldEmptyString() {
assertEquals("", CwCtcDecoder.greedy(logits(intArrayOf(blank, blank, blank)), chars, blank))
}
@Test
fun greedy_emptyInputYieldsEmptyString() {
assertEquals("", CwCtcDecoder.greedy(logits(intArrayOf()), chars, blank))
}
@Test
fun greedy_decodesCallsignWithSpaceAndDigits() {
// "CQ BG7" — C=16 Q=30 space=40 B=15 G=20 7=10
val path = intArrayOf(
blank, 16, 16, blank, 30, blank, 40,
15, blank, 20, blank, 10, blank
)
assertEquals("CQ BG7", CwCtcDecoder.greedy(logits(path), chars, blank))
}
@Test
fun greedy_picksHighestScoringClassPerFrame() {
// Frame favours S (32) over T (33); only S must survive.
val frame = FloatArray(42) { -10f }
frame[33] = -1f
frame[32] = -0.1f
assertEquals("S", CwCtcDecoder.greedy(arrayOf(arrayOf(frame)), chars, blank))
}
}