From 24aec7597b10c464c2f6b79f706c7cead199a796 Mon Sep 17 00:00:00 2001 From: idevlab Date: Fri, 18 Sep 2026 12:33:43 +0800 Subject: [PATCH] feat: inject vocabulary hotwords into Qwen3-ASR and Volc ASR MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Volc: build request.corpus.context as the documented {hotwords:[{word:...}]} JSON string from the recognition phrases, bounded for the 双向流式 direct-pass budget, and include it in the full client request. The streaming partials and the final file transcription share the same request path. Qwen3-ASR: mlx-audio-swift exposes Qwen3ASRModel.generate(context:) which injects the text into the system prompt. Add configureRecognition storage and pass a bounded terms prompt; the generic STTGenerationModel protocol has no context parameter, so no other engine is touched. Co-authored-by: multica-agent --- Sources/Speech/QwenNativeASREngine.swift | 29 ++++++- Sources/Speech/SpeechRecognitionContext.swift | 19 +++++ Sources/Speech/VolcSpeechEngine.swift | 85 ++++++++++++++++--- .../QwenNativeASREngineTests.swift | 13 +++ .../SpeechRecognitionQualityTests.swift | 14 +++ .../VolcSpeechEnginePayloadTests.swift | 62 ++++++++++++++ 6 files changed, 207 insertions(+), 15 deletions(-) create mode 100644 Tests/OpenTypeTests/VolcSpeechEnginePayloadTests.swift diff --git a/Sources/Speech/QwenNativeASREngine.swift b/Sources/Speech/QwenNativeASREngine.swift index b4bae3ec..b243b8a3 100644 --- a/Sources/Speech/QwenNativeASREngine.swift +++ b/Sources/Speech/QwenNativeASREngine.swift @@ -5,6 +5,8 @@ import MLXAudioSTT final class QwenNativeASREngine: SpeechEngine, @unchecked Sendable { private let modelDirectory: URL private let runtime = QwenNativeASRRuntime() + private let recognitionContextLock = NSLock() + private var recognitionContext = SpeechRecognitionContext.empty init(modelPath: String) { modelDirectory = URL(fileURLWithPath: modelPath).standardizedFileURL @@ -18,6 +20,20 @@ final class QwenNativeASREngine: SpeechEngine, @unchecked Sendable { modelDirectory == URL(fileURLWithPath: modelPath).standardizedFileURL } + func configureRecognition(context: SpeechRecognitionContext) { + recognitionContextLock.lock() + recognitionContext = context + recognitionContextLock.unlock() + } + + /// The exact context string handed to `Qwen3ASRModel.generate(context:)`. + /// Exposed for tests so vocabulary injection can be asserted without a model. + func currentContextPrompt() -> String? { + recognitionContextLock.lock() + defer { recognitionContextLock.unlock() } + return recognitionContext.contextualPrompt() + } + func prepare() async { guard isReady else { return } do { @@ -31,12 +47,14 @@ final class QwenNativeASREngine: SpeechEngine, @unchecked Sendable { guard isReady else { throw QwenNativeASRError.notConfigured } guard let audioURL else { throw QwenNativeASRError.noAudioFile } + let contextPrompt = currentContextPrompt() let started = CFAbsoluteTimeGetCurrent() let result = try await QwenAudioPreprocessor.withPreparedAudio(from: audioURL) { preparedURL in try await runtime.transcribe( audioURL: preparedURL, modelDirectory: modelDirectory, - language: language + language: language, + context: contextPrompt ) } let elapsed = CFAbsoluteTimeGetCurrent() - started @@ -92,7 +110,8 @@ private actor QwenNativeASRRuntime { func transcribe( audioURL: URL, modelDirectory: URL, - language: String? + language: String?, + context: String? ) async throws -> Result { try Task.checkCancellation() let model = try await loadModel(from: modelDirectory) @@ -104,7 +123,11 @@ private actor QwenNativeASRRuntime { throw QwenAudioPreprocessorError.conversionFailed } - let output = model.generate(audio: audio, language: language) + let output = model.generate( + audio: audio, + context: context ?? "", + language: language + ) try Task.checkCancellation() return Result( text: output.text, diff --git a/Sources/Speech/SpeechRecognitionContext.swift b/Sources/Speech/SpeechRecognitionContext.swift index a599c9a0..32538e5f 100644 --- a/Sources/Speech/SpeechRecognitionContext.swift +++ b/Sources/Speech/SpeechRecognitionContext.swift @@ -92,4 +92,23 @@ struct SpeechRecognitionContext: Equatable, Sendable { } return bestTokens } + + /// Free-text biasing prompt for engines whose API accepts a text context + /// (e.g. Qwen3-ASR `context`, injected into the system prompt). Terms are + /// added in rank order until the character budget is reached, so a large + /// dictionary cannot grow the prompt without bound. + func contextualPrompt(maximumCharacters: Int = 600) -> String? { + guard maximumCharacters > 0 else { return nil } + let prefix = "Terms: " + var accepted: [String] = [] + var characterCount = prefix.count + for phrase in phrases { + let addition = (accepted.isEmpty ? 0 : 2) + phrase.count + guard characterCount + addition <= maximumCharacters else { continue } + accepted.append(phrase) + characterCount += addition + } + guard !accepted.isEmpty else { return nil } + return prefix + accepted.joined(separator: ", ") + } } diff --git a/Sources/Speech/VolcSpeechEngine.swift b/Sources/Speech/VolcSpeechEngine.swift index ad0d6c6b..4d054c38 100644 --- a/Sources/Speech/VolcSpeechEngine.swift +++ b/Sources/Speech/VolcSpeechEngine.swift @@ -9,10 +9,16 @@ final class VolcSpeechEngine: SpeechEngine, @unchecked Sendable { private(set) var isReady: Bool private typealias Connection = (session: URLSession, task: URLSessionWebSocketTask) private var streamingSession: VolcStreamingSession? + private let recognitionContextLock = NSLock() + private var recognitionContext = SpeechRecognitionContext.empty private static let endpoint = "wss://openspeech.bytedance.com/api/v3/sauc/bigmodel" private static let chunkSize = 6400 // ~200ms at 16kHz 16-bit mono private static let timeoutSeconds: UInt64 = 30 + /// The 双向流式 endpoint's direct hotword context budget is documented as + /// 100 tokens, so keep the list short and bounded well below that. + static let maximumHotwordCount = 100 + static let maximumHotwordCharacters = 300 init(appKey: String, accessKey: String, resourceId: String) { self.appKey = appKey @@ -23,6 +29,18 @@ final class VolcSpeechEngine: SpeechEngine, @unchecked Sendable { var supportsStreaming: Bool { true } + func configureRecognition(context: SpeechRecognitionContext) { + recognitionContextLock.lock() + recognitionContext = context + recognitionContextLock.unlock() + } + + private func recognitionContextSnapshot() -> SpeechRecognitionContext { + recognitionContextLock.lock() + defer { recognitionContextLock.unlock() } + return recognitionContext + } + func startListening(language: String?, onPartialResult: @escaping @Sendable (String) -> Void) { guard isReady else { return } streamingSession = VolcStreamingSession( @@ -140,6 +158,48 @@ final class VolcSpeechEngine: SpeechEngine, @unchecked Sendable { // MARK: - Send full client request private func sendFullClientRequest(conn: Connection, language: String?) async throws { + let hotwords = Self.hotwordContext( + for: recognitionContextSnapshot().phrases + ) + let payload = Self.fullClientRequestPayload( + language: language, + hotwordContext: hotwords + ) + + let jsonData = try JSONSerialization.data(withJSONObject: payload) + let message = buildMessage(type: .fullClientRequest, flags: 0x00, serialization: .json, payload: jsonData) + try await sendMessage(conn: conn, data: message) + } + + /// Builds the `request.corpus.context` hotword payload. Volc accepts a JSON + /// string of `{"hotwords":[{"word": ...}]}`; see the official + /// 大模型流式语音识别 API `corpus.context` field. + static func hotwordContext(for phrases: [String]) -> String? { + var accepted: [String] = [] + var characterCount = 0 + for phrase in phrases { + let word = phrase.trimmingCharacters(in: .whitespacesAndNewlines) + guard !word.isEmpty, accepted.count < maximumHotwordCount else { continue } + guard characterCount + word.count <= maximumHotwordCharacters else { continue } + accepted.append(word) + characterCount += word.count + } + guard !accepted.isEmpty else { return nil } + + let payload: [String: Any] = ["hotwords": accepted.map { ["word": $0] }] + guard let data = try? JSONSerialization.data( + withJSONObject: payload, + options: [.sortedKeys] + ) else { + return nil + } + return String(data: data, encoding: .utf8) + } + + static func fullClientRequestPayload( + language: String?, + hotwordContext: String? + ) -> [String: Any] { var audio: [String: Any] = [ "format": "pcm", "rate": 16000, @@ -147,22 +207,23 @@ final class VolcSpeechEngine: SpeechEngine, @unchecked Sendable { "channel": 1, "codec": "raw" ] - if let lang = language { audio["language"] = lang } + if let language { audio["language"] = language } - let payload: [String: Any] = [ + var request: [String: Any] = [ + "model_name": "bigmodel", + "enable_itn": true, + "enable_punc": true, + "show_utterances": false + ] + if let hotwordContext, !hotwordContext.isEmpty { + request["corpus"] = ["context": hotwordContext] + } + + return [ "user": ["uid": "opentype_macos"], "audio": audio, - "request": [ - "model_name": "bigmodel", - "enable_itn": true, - "enable_punc": true, - "show_utterances": false - ] + "request": request ] - - let jsonData = try JSONSerialization.data(withJSONObject: payload) - let message = buildMessage(type: .fullClientRequest, flags: 0x00, serialization: .json, payload: jsonData) - try await sendMessage(conn: conn, data: message) } // MARK: - Stream audio diff --git a/Tests/OpenTypeTests/QwenNativeASREngineTests.swift b/Tests/OpenTypeTests/QwenNativeASREngineTests.swift index 93a18361..fecc47ed 100644 --- a/Tests/OpenTypeTests/QwenNativeASREngineTests.swift +++ b/Tests/OpenTypeTests/QwenNativeASREngineTests.swift @@ -112,6 +112,19 @@ final class QwenNativeASREngineTests: XCTestCase { } } + func testRecognitionContextPromptReachesTheModelCall() { + let engine = QwenNativeASREngine(modelPath: "/nonexistent-qwen-model") + XCTAssertNil(engine.currentContextPrompt()) + + engine.configureRecognition( + context: SpeechRecognitionContext(phrases: ["OpenType", "菜单栏"]) + ) + XCTAssertEqual(engine.currentContextPrompt(), "Terms: OpenType, 菜单栏") + + engine.configureRecognition(context: .empty) + XCTAssertNil(engine.currentContextPrompt()) + } + private func safetensorsFiles(in directory: URL) throws -> Set { let files = try FileManager.default.contentsOfDirectory( at: directory, diff --git a/Tests/OpenTypeTests/SpeechRecognitionQualityTests.swift b/Tests/OpenTypeTests/SpeechRecognitionQualityTests.swift index ae179677..92d351ef 100644 --- a/Tests/OpenTypeTests/SpeechRecognitionQualityTests.swift +++ b/Tests/OpenTypeTests/SpeechRecognitionQualityTests.swift @@ -121,6 +121,20 @@ final class SpeechRecognitionQualityTests: XCTestCase { XCTAssertEqual(prompt, "Dictation terms: OpenType, MLX.") } + func testContextualPromptListsTermsAndRespectsBudget() { + XCTAssertNil(SpeechRecognitionContext(phrases: []).contextualPrompt()) + + let context = SpeechRecognitionContext(phrases: ["OpenType", "菜单栏"]) + XCTAssertEqual(context.contextualPrompt(), "Terms: OpenType, 菜单栏") + XCTAssertEqual(context.contextualPrompt(maximumCharacters: 15), "Terms: OpenType") + XCTAssertNil(context.contextualPrompt(maximumCharacters: 9)) + + let longFirst = SpeechRecognitionContext( + phrases: [String(repeating: "x", count: 80), "MLX"] + ) + XCTAssertEqual(longFirst.contextualPrompt(maximumCharacters: 20), "Terms: MLX") + } + func testAppleCompatibleDictationPresetChangesAfterOneMinute() { XCTAssertEqual( AppleSpeechAnalyzer.dictationPreset(forDuration: 30), diff --git a/Tests/OpenTypeTests/VolcSpeechEnginePayloadTests.swift b/Tests/OpenTypeTests/VolcSpeechEnginePayloadTests.swift new file mode 100644 index 00000000..6aad8eb8 --- /dev/null +++ b/Tests/OpenTypeTests/VolcSpeechEnginePayloadTests.swift @@ -0,0 +1,62 @@ +import XCTest +@testable import OpenType + +final class VolcSpeechEnginePayloadTests: XCTestCase { + func testHotwordContextSerializesPhrasesAsWordEntries() throws { + let context = try XCTUnwrap( + VolcSpeechEngine.hotwordContext(for: ["OpenType", "菜单栏"]) + ) + + let object = try XCTUnwrap( + JSONSerialization.jsonObject(with: Data(context.utf8)) as? [String: Any] + ) + let hotwords = try XCTUnwrap(object["hotwords"] as? [[String: String]]) + + XCTAssertEqual(hotwords.compactMap { $0["word"] }, ["OpenType", "菜单栏"]) + } + + func testFullClientRequestCarriesHotwordsInCorpusContext() throws { + let phrases = ["OpenType", "菜单栏"] + let context = try XCTUnwrap(VolcSpeechEngine.hotwordContext(for: phrases)) + + let payload = VolcSpeechEngine.fullClientRequestPayload( + language: "zh-CN", + hotwordContext: context + ) + + let request = try XCTUnwrap(payload["request"] as? [String: Any]) + XCTAssertEqual(request["model_name"] as? String, "bigmodel") + let corpus = try XCTUnwrap(request["corpus"] as? [String: Any]) + XCTAssertEqual(corpus["context"] as? String, context) + + let audio = try XCTUnwrap(payload["audio"] as? [String: Any]) + XCTAssertEqual(audio["language"] as? String, "zh-CN") + } + + func testPayloadOmitsCorpusWithoutPhrases() { + XCTAssertNil(VolcSpeechEngine.hotwordContext(for: [])) + XCTAssertNil(VolcSpeechEngine.hotwordContext(for: [" "])) + + let payload = VolcSpeechEngine.fullClientRequestPayload( + language: nil, + hotwordContext: nil + ) + let request = payload["request"] as? [String: Any] + XCTAssertNil(request?["corpus"]) + } + + func testHotwordBudgetKeepsWholeEntriesOnly() throws { + let phrases = (0..<150).map { "term\($0)" } + let context = try XCTUnwrap(VolcSpeechEngine.hotwordContext(for: phrases)) + + let object = try XCTUnwrap( + JSONSerialization.jsonObject(with: Data(context.utf8)) as? [String: Any] + ) + let hotwords = try XCTUnwrap(object["hotwords"] as? [[String: String]]) + + XCTAssertLessThanOrEqual(hotwords.count, VolcSpeechEngine.maximumHotwordCount) + for entry in hotwords { + XCTAssertTrue(phrases.contains(entry["word"] ?? ""), entry["word"] ?? "") + } + } +}