Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 26 additions & 3 deletions Sources/Speech/QwenNativeASREngine.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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 {
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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,
Expand Down
19 changes: 19 additions & 0 deletions Sources/Speech/SpeechRecognitionContext.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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: ", ")
}
}
85 changes: 73 additions & 12 deletions Sources/Speech/VolcSpeechEngine.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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(
Expand Down Expand Up @@ -140,29 +158,72 @@ 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,
"bits": 16,
"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
Expand Down
13 changes: 13 additions & 0 deletions Tests/OpenTypeTests/QwenNativeASREngineTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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<String> {
let files = try FileManager.default.contentsOfDirectory(
at: directory,
Expand Down
14 changes: 14 additions & 0 deletions Tests/OpenTypeTests/SpeechRecognitionQualityTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
62 changes: 62 additions & 0 deletions Tests/OpenTypeTests/VolcSpeechEnginePayloadTests.swift
Original file line number Diff line number Diff line change
@@ -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"] ?? "")
}
}
}
Loading