diff --git a/Documentation/ASR/CustomVocabulary.md b/Documentation/ASR/CustomVocabulary.md index 411a881fa..a9c5b6450 100644 --- a/Documentation/ASR/CustomVocabulary.md +++ b/Documentation/ASR/CustomVocabulary.md @@ -24,6 +24,43 @@ The paper introduces a dynamic programming algorithm for CTC-based keyword spott Both approaches achieve identical 99.4% accuracy. Approach 1 is faster but only works with TDT-CTC-110M because that model has a built-in CTC head. Approach 2 works with any TDT model but loads a separate CTC encoder. +## Optional preflight for transcript-only consumers + +A prepared `VocabularyBoostingSession` can check whether a transcript has any +candidate requiring CTC evidence before running acoustic rescoring: + +```swift +if session.hasCTCRescoringCandidates( + text: result.text, + tokenTimings: result.tokenTimings ?? [] +) { + let rescored = await session.rescore( + text: result.text, + tokenTimings: result.tokenTimings ?? [], + audioSamples: audio + ) + // Use rescored.text when a replacement was applied. +} +// Otherwise keep result.text. +``` + +This uses the same aliases, compounds, similarity thresholds and safety rules as +rescoring. It performs no CTC inference. An empty vocabulary or timing list +returns `false`. When acoustic rescue is enabled on the term-centric path and the +vocabulary size is at or below `ContextBiasingConstants.largeVocabThreshold`, +nonempty input returns `true` conservatively: the spotter can find a term even +when text matching cannot. Larger vocabularies and the experimental BK-tree path +use text candidate discovery because acoustic rescue does not run there. +Disabling rescue changes recognition behavior; choose that policy separately +from whether to use preflight. + +Use this only when consuming transcript replacements. Applications that also +need `detectedTerms` must still run `rescore`: keyword detections can exist +without a text replacement candidate. The existing `rescore` API retains that +behavior. Callers managing their own CTC head can use +`VocabularyRescorer.hasCTCRescoringCandidates` with the same `minSimilarity` they +pass to `ctcTokenRescore`. + ## Model Compatibility FluidAudio supports two ASR models with different architectures: diff --git a/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer+TokenRescoring.swift b/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer+TokenRescoring.swift index b234679ac..15a033348 100644 --- a/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer+TokenRescoring.swift +++ b/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer+TokenRescoring.swift @@ -346,6 +346,7 @@ extension VocabularyRescorer { minSimilarity: Float = ContextBiasingConstants.minSimilarityFloor ) -> RescoreOutput { var candidateEvidence: CandidateEvidenceCollector? + var candidateFound: Bool? return evaluateTokenCandidates( transcript: transcript, tokenTimings: tokenTimings, @@ -354,7 +355,8 @@ extension VocabularyRescorer { cbw: cbw, marginSeconds: marginSeconds, minSimilarity: minSimilarity, - candidateEvidence: &candidateEvidence + candidateEvidence: &candidateEvidence, + candidateFound: &candidateFound ) } @@ -385,6 +387,7 @@ extension VocabularyRescorer { minSimilarity: Float = ContextBiasingConstants.minSimilarityFloor ) -> CandidateEvidenceOutput { var candidateEvidence: CandidateEvidenceCollector? = CandidateEvidenceCollector() + var candidateFound: Bool? _ = evaluateTokenCandidates( transcript: transcript, tokenTimings: tokenTimings, @@ -393,7 +396,8 @@ extension VocabularyRescorer { cbw: cbw, marginSeconds: marginSeconds, minSimilarity: minSimilarity, - candidateEvidence: &candidateEvidence + candidateEvidence: &candidateEvidence, + candidateFound: &candidateFound ) return CandidateEvidenceOutput( baseText: transcript, @@ -402,6 +406,47 @@ extension VocabularyRescorer { ) } + /// Whether rescoring can require acoustic evidence for this transcript. + /// + /// Uses the same candidate discovery and safety rules as ``ctcTokenRescore`` + /// without computing or consuming CTC probabilities. A false result permits + /// skipping acoustic inference for transcript replacements. Callers needing + /// standalone keyword detections must still run acoustic inference. When + /// term-centric acoustic rescue is enabled for this vocabulary size, returns + /// true conservatively because rescue can find terms absent from text candidates. + /// + /// - Parameters: + /// - transcript: Untouched transcript from the TDT decoder. + /// - tokenTimings: Token-level timings from the TDT decoder. + /// - minSimilarity: The same vocabulary threshold used for rescoring. + /// - Returns: False only when rescoring has no candidate requiring CTC evidence. + public func hasCTCRescoringCandidates( + transcript: String, + tokenTimings: [TokenTiming], + minSimilarity: Float = ContextBiasingConstants.minSimilarityFloor + ) -> Bool { + guard !vocabulary.terms.isEmpty, !tokenTimings.isEmpty else { return false } + if config.spotterRescueEnabled, !useBKTree, + vocabulary.terms.count <= ContextBiasingConstants.largeVocabThreshold + { + return true + } + var candidateEvidence: CandidateEvidenceCollector? + var candidateFound: Bool? = false + _ = evaluateTokenCandidates( + transcript: transcript, + tokenTimings: tokenTimings, + logProbs: [], + frameDuration: 0, + cbw: 0, + marginSeconds: 0, + minSimilarity: minSimilarity, + candidateEvidence: &candidateEvidence, + candidateFound: &candidateFound + ) + return candidateFound == true + } + private func evaluateTokenCandidates( transcript: String, tokenTimings: [TokenTiming], @@ -410,7 +455,8 @@ extension VocabularyRescorer { cbw: Float, marginSeconds: Double, minSimilarity: Float, - candidateEvidence: inout CandidateEvidenceCollector? + candidateEvidence: inout CandidateEvidenceCollector?, + candidateFound: inout Bool? ) -> RescoreOutput { // Build word-level timings once at the entrypoint and pass into both // dispatch paths. Computing this once instead of twice avoids @@ -438,7 +484,8 @@ extension VocabularyRescorer { cbw: cbw, marginSeconds: marginSeconds, minSimilarity: minSimilarity, - candidateEvidence: &candidateEvidence + candidateEvidence: &candidateEvidence, + candidateFound: &candidateFound ) } else { return rescoreWithConstrainedCTCTermCentric( @@ -449,7 +496,8 @@ extension VocabularyRescorer { cbw: cbw, marginSeconds: marginSeconds, minSimilarity: minSimilarity, - candidateEvidence: &candidateEvidence + candidateEvidence: &candidateEvidence, + candidateFound: &candidateFound ) } } @@ -472,9 +520,10 @@ extension VocabularyRescorer { cbw: Float = ContextBiasingConstants.defaultCbw, marginSeconds: Double = ContextBiasingConstants.defaultMarginSeconds, minSimilarity: Float = ContextBiasingConstants.minSimilarityFloor, - candidateEvidence: inout CandidateEvidenceCollector? + candidateEvidence: inout CandidateEvidenceCollector?, + candidateFound: inout Bool? ) -> RescoreOutput { - guard !wordTimings.isEmpty, !logProbs.isEmpty else { + guard !wordTimings.isEmpty, candidateFound != nil || !logProbs.isEmpty else { return RescoreOutput(text: transcript, replacements: [], wasModified: false) } @@ -491,9 +540,6 @@ extension VocabularyRescorer { var replacedIndices = Set() var pendingReplacements: [PendingReplacement] = [] - // Build normalized vocabulary set for guard checks - let vocabularyNormalizedSet = buildVocabularyNormalizedSet() - // Lowest per-term similarity across the vocabulary. The BK-tree search // bound is derived from this floor so that terms with a lower per-term // `minSimilarity` are not pruned before per-candidate filtering applies @@ -644,6 +690,11 @@ extension VocabularyRescorer { spanEndTime: spanEndTime ) + // Stop at the exact scoring boundary, after all text safety guards. + if candidateFound != nil { + candidateFound = true + return RescoreOutput(text: transcript, replacements: [], wasModified: false) + } let result = evaluateCTCMatch( candidate: matchCandidate, logProbs: logProbs, @@ -698,9 +749,10 @@ extension VocabularyRescorer { cbw: Float = ContextBiasingConstants.defaultCbw, marginSeconds: Double = ContextBiasingConstants.defaultMarginSeconds, minSimilarity: Float = ContextBiasingConstants.minSimilarityFloor, - candidateEvidence: inout CandidateEvidenceCollector? + candidateEvidence: inout CandidateEvidenceCollector?, + candidateFound: inout Bool? ) -> RescoreOutput { - guard !wordTimings.isEmpty, !logProbs.isEmpty else { + guard !wordTimings.isEmpty, candidateFound != nil || !logProbs.isEmpty else { return RescoreOutput(text: transcript, replacements: [], wasModified: false) } @@ -716,8 +768,7 @@ extension VocabularyRescorer { var replacedIndices = Set() var pendingReplacements: [PendingReplacement] = [] // Two-pass: collect first, apply later - // Build normalized vocabulary set for guard checks - let vocabularyNormalizedSet = buildVocabularyNormalizedSet() + let normalizedWords = wordTimings.map { Self.normalizeForSimilarity($0.word) } // TERM-CENTRIC LOOP: For each vocabulary term, find similar TDT words and run constrained CTC for term in vocabulary.terms { @@ -826,6 +877,11 @@ extension VocabularyRescorer { spanEndTime: spanEndTime ) + // Stop at the exact scoring boundary, after all text safety guards. + if candidateFound != nil { + candidateFound = true + return RescoreOutput(text: transcript, replacements: [], wasModified: false) + } let result = evaluateCTCMatch( candidate: matchCandidate, logProbs: logProbs, @@ -860,7 +916,7 @@ extension VocabularyRescorer { guard !replacedIndices.contains(wordIdx) else { continue } let tdtWord = timing.word - let normalizedWord = Self.normalizeForSimilarity(tdtWord) + let normalizedWord = normalizedWords[wordIdx] guard !normalizedWord.isEmpty else { continue } // Skip if already exact match to canonical (no replacement needed) @@ -898,11 +954,11 @@ extension VocabularyRescorer { // Pre-compute normalized adjacent words (only if needed) let normalized2: String? = (wordIdx + 1 < wordTimings.count && !replacedIndices.contains(wordIdx + 1)) - ? Self.normalizeForSimilarity(wordTimings[wordIdx + 1].word) + ? normalizedWords[wordIdx + 1] : nil let normalized3: String? = (wordIdx + 2 < wordTimings.count && !replacedIndices.contains(wordIdx + 2)) - ? Self.normalizeForSimilarity(wordTimings[wordIdx + 2].word) + ? normalizedWords[wordIdx + 2] : nil // 2-word compound matching @@ -965,7 +1021,7 @@ extension VocabularyRescorer { // STOPWORD CHECKS let spanWords = matchedSpanLength >= 2 - ? (0.. [TermFormKey: [NormalizedForm]] { + var aliasesByCanonical: [String: [String]] = [:] + for term in vocabulary.terms { + aliasesByCanonical[term.textLowercased, default: []].append(contentsOf: term.aliases ?? []) + } + var formsByText: [TermFormKey: [NormalizedForm]] = [:] + for term in vocabulary.terms { + formsByText[TermFormKey(term)] = Self.normalizedForms( + canonicalTerm: term.text, + aliases: (aliasesByCanonical[term.textLowercased] ?? []) + (term.aliases ?? [])) + } + return formsByText + } + /// Build all normalized forms (canonical + aliases) for a vocabulary term func buildNormalizedForms(for term: CustomVocabularyTerm) -> [NormalizedForm] { + if let prepared = normalizedTermForms[TermFormKey(term)] { return prepared } var aliases: [String] = [] let termLower = term.textLowercased @@ -345,26 +361,6 @@ extension VocabularyRescorer { "#", ".", "@", "%", "&", "*", "/", "\\", "_", "-", "`", "'", "’", "^", ] - /// Build set of normalized vocabulary terms for guard checks - func buildVocabularyNormalizedSet() -> Set { - var normalizedSet = Set() - for term in vocabulary.terms { - let normalized = Self.normalizeForSimilarity(term.text) - if !normalized.isEmpty { - normalizedSet.insert(normalized) - } - // Also add aliases if present - if let aliases = term.aliases { - for alias in aliases { - let normalizedAlias = Self.normalizeForSimilarity(alias) - if !normalizedAlias.isEmpty { - normalizedSet.insert(normalizedAlias) - } - } - } - } - return normalizedSet - } } // MARK: - Token Word Boundary Utilities diff --git a/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer.swift b/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer.swift index b24e2d69f..d38164e71 100644 --- a/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer.swift +++ b/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/Rescorer/VocabularyRescorer.swift @@ -17,6 +17,18 @@ public struct VocabularyRescorer: Sendable { let vocabulary: CustomVocabularyContext let ctcTokenizer: CtcTokenizer? let debugMode: Bool + let normalizedTermForms: [TermFormKey: [NormalizedForm]] + let vocabularyNormalizedSet: Set + + struct TermFormKey: Hashable, Sendable { + let text: String + let aliases: [String] + + init(_ term: CustomVocabularyTerm) { + text = term.text + aliases = term.aliases ?? [] + } + } // BK-tree for efficient approximate string matching (experimental) // When enabled, uses BK-tree to find candidate vocabulary terms within edit distance @@ -180,6 +192,9 @@ public struct VocabularyRescorer: Sendable { bkTree: BKTree?, bkTreeMaxDistance: Int ) { + let formsByText = Self.prepareNormalizedForms(for: vocabulary) + self.normalizedTermForms = formsByText + self.vocabularyNormalizedSet = Set(formsByText.values.flatMap { $0.map(\.normalized) }) self.spotter = spotter self.vocabulary = vocabulary self.config = config diff --git a/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/VocabularyBoostingSession.swift b/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/VocabularyBoostingSession.swift index b968243ae..84f238324 100644 --- a/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/VocabularyBoostingSession.swift +++ b/Sources/FluidAudio/ASR/Parakeet/SlidingWindow/CustomVocabulary/VocabularyBoostingSession.swift @@ -84,6 +84,16 @@ public struct VocabularyBoostingSession: Sendable { ) } + /// Whether transcript rescoring needs CTC inference. This does not predict + /// standalone keyword detections; callers requiring those must still rescore. + public func hasCTCRescoringCandidates(text: String, tokenTimings: [TokenTiming]) -> Bool { + rescorer.hasCTCRescoringCandidates( + transcript: text, + tokenTimings: tokenTimings, + minSimilarity: max(vocabSizeConfig.minSimilarity, vocabulary.minSimilarity) + ) + } + /// Rescore a transcript against CTC acoustic evidence from its audio. /// /// `tokenTimings` must be on the same clock as `audioSamples`: time zero diff --git a/Tests/FluidAudioTests/ASR/Parakeet/SlidingWindow/CustomVocabulary/VocabularyCandidateDiscoveryTests.swift b/Tests/FluidAudioTests/ASR/Parakeet/SlidingWindow/CustomVocabulary/VocabularyCandidateDiscoveryTests.swift new file mode 100644 index 000000000..e4cd3a2f5 --- /dev/null +++ b/Tests/FluidAudioTests/ASR/Parakeet/SlidingWindow/CustomVocabulary/VocabularyCandidateDiscoveryTests.swift @@ -0,0 +1,171 @@ +import AVFoundation +import Foundation +import XCTest + +@testable import FluidAudio + +@MainActor +final class VocabularyCandidateDiscoveryTests: XCTestCase { + func testPreparedFormsPreserveAliasOrderingAndGuardSetWithoutModels() throws { + let terms = [ + CustomVocabularyTerm(text: "ESLint", aliases: ["E S lint", "es lint"]), + CustomVocabularyTerm(text: "eslint", aliases: ["E S LINT", "easylint"]), + CustomVocabularyTerm(text: "Claude Code", aliases: ["cloud code", "", "!!!"]), + CustomVocabularyTerm(text: "AI", aliases: nil), + ] + let forms = VocabularyRescorer.prepareNormalizedForms(for: CustomVocabularyContext(terms: terms)) + for term in terms { + let aliases = + terms.filter { $0.textLowercased == term.textLowercased } + .flatMap { $0.aliases ?? [] } + (term.aliases ?? []) + let expected = VocabularyRescorer.normalizedForms(canonicalTerm: term.text, aliases: aliases) + XCTAssertEqual(try XCTUnwrap(forms[VocabularyRescorer.TermFormKey(term)]), expected) + } + let originalGuardSet = Set( + terms.flatMap { [$0.text] + ($0.aliases ?? []) } + .map(VocabularyRescorer.normalizeForSimilarity).filter { !$0.isEmpty }) + XCTAssertEqual(Set(forms.values.flatMap { $0.map(\.normalized) }), originalGuardSet) + XCTAssertTrue(VocabularyRescorer.prepareNormalizedForms(for: CustomVocabularyContext(terms: [])).isEmpty) + } + + func testRescuePreflightRespectsVocabularySizeWithInstalledModels() async throws { + let directory = CtcModels.defaultCacheDirectory(for: .ctc110m) + guard FileManager.default.fileExists(atPath: directory.appendingPathComponent("tokenizer.json").path) else { + throw XCTSkip("Install CTC 110M to exercise rescue preflight") + } + let models = try await CtcModels.loadDirect(from: directory) + let spotter = CtcKeywordSpotter(models: models) + let tokenizer = try await CtcTokenizer.load(from: directory) + let threshold = ContextBiasingConstants.largeVocabThreshold + let unrelatedTimings = [TokenTiming(token: "▁unrelated", tokenId: 1, startTime: 0, endTime: 1, confidence: 1)] + let candidateTimings = [TokenTiming(token: "▁quiltor", tokenId: 1, startTime: 0, endTime: 1, confidence: 1)] + for count in [threshold, threshold + 1] { + let terms = ["Quilter"] + (1.. \(term)") + let session = try await VocabularyBoostingSession( + vocabulary: context, ctcModels: models, config: .init(spotterRescueEnabled: false)) + let sessionThreshold = ContextBiasingConstants.rescorerConfig(forVocabSize: context.terms.count) + .minSimilarity + XCTAssertEqual( + session.hasCTCRescoringCandidates(text: transcript, tokenTimings: timings), + rescorer.hasCTCRescoringCandidates( + transcript: transcript, tokenTimings: timings, + minSimilarity: max(sessionThreshold, context.minSimilarity))) + if discovery { positive += 1 } else { negative += 1 } + XCTAssertFalse(rescorer.hasCTCRescoringCandidates(transcript: transcript, tokenTimings: [])) + if threshold == 0.99 { XCTAssertFalse(discovery) } + } + XCTAssertGreaterThan(positive, 0) + XCTAssertGreaterThan(negative, 0) + let rescue = try await VocabularyRescorer.create( + spotter: spotter, + vocabulary: CustomVocabularyContext(terms: [CustomVocabularyTerm(text: "Quilter")]), + ctcModelDirectory: directory) + XCTAssertFalse(rescue.hasCTCRescoringCandidates(transcript: "", tokenTimings: [])) + let unrelatedTimings = [TokenTiming(token: "▁unrelated", tokenId: 1, startTime: 0, endTime: 1, confidence: 1)] + XCTAssertTrue( + rescue.hasCTCRescoringCandidates(transcript: "unrelated", tokenTimings: unrelatedTimings), + "Acoustic rescue can require evidence without a text candidate") + let empty = try await VocabularyRescorer.create( + spotter: spotter, vocabulary: CustomVocabularyContext(terms: []), ctcModelDirectory: directory) + XCTAssertFalse(empty.hasCTCRescoringCandidates(transcript: "unrelated", tokenTimings: unrelatedTimings)) + } +}