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
39 changes: 39 additions & 0 deletions Sources/Speech/TranscriptSegmentJoiner.swift
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
import Foundation

/// Joins the per-chunk texts WhisperKit returns when VAD chunking splits a
/// recording. Latin scripts need a separating space; CJK scripts must not get
/// one. When the language is auto-detected (`nil`), fall back to inspecting the
/// boundary characters so CJK output still joins without spaces.
enum TranscriptSegmentJoiner {
static func joined(_ segments: [String], language: String?) -> String {
let cleaned = segments
.map { $0.trimmingCharacters(in: .whitespacesAndNewlines) }
.filter { !$0.isEmpty }
guard var result = cleaned.first else { return "" }

for segment in cleaned.dropFirst() {
result += separator(previous: result, next: segment, language: language)
result += segment
}
return result
}

static func separator(previous: String, next: String, language: String?) -> String {
if let language, !language.isEmpty {
return usesNoSpaceScript(language) ? "" : " "
}
if let last = previous.last, let first = next.first,
isNoSpaceBoundary(last) || isNoSpaceBoundary(first) {
return ""
}
return " "
}

private static func usesNoSpaceScript(_ language: String) -> Bool {
["zh", "yue", "ja", "ko"].contains(language.lowercased())
}

private static func isNoSpaceBoundary(_ character: Character) -> Bool {
character.unicodeScalars.contains(where: NoSpaceScript.contains)
}
}
9 changes: 5 additions & 4 deletions Sources/Speech/WhisperEngine.swift
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,7 @@ final class WhisperEngine: SpeechEngine, @unchecked Sendable {
)
streamingSession = WhisperStreamingSession(
whisperKit: whisperKit,
language: language,
partialHandler: onPartialResult,
optionsBuilder: { options }
)
Expand Down Expand Up @@ -164,10 +165,10 @@ final class WhisperEngine: SpeechEngine, @unchecked Sendable {
audioPath: url.path,
decodeOptions: options
)
let text = results
.compactMap { $0.text }
.joined(separator: " ")
.trimmingCharacters(in: .whitespacesAndNewlines)
let text = TranscriptSegmentJoiner.joined(
results.compactMap { $0.text },
language: language
)

let elapsed = CFAbsoluteTimeGetCurrent() - t0
Log.info("[WhisperEngine] transcribed \(text.count) chars in \(String(format: "%.1f", elapsed))s")
Expand Down
11 changes: 7 additions & 4 deletions Sources/Speech/WhisperStreamingSession.swift
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import WhisperKit

final class WhisperStreamingSession: @unchecked Sendable {
private let whisperKit: WhisperKit
private let language: String?
private let partialHandler: @Sendable (String) -> Void
private let optionsBuilder: () -> DecodingOptions
private let queue = DispatchQueue(label: "opentype.whisper-stream")
Expand All @@ -23,10 +24,12 @@ final class WhisperStreamingSession: @unchecked Sendable {

init(
whisperKit: WhisperKit,
language: String? = nil,
partialHandler: @escaping @Sendable (String) -> Void,
optionsBuilder: @escaping () -> DecodingOptions
) {
self.whisperKit = whisperKit
self.language = language
self.partialHandler = partialHandler
self.optionsBuilder = optionsBuilder
}
Expand Down Expand Up @@ -134,10 +137,10 @@ final class WhisperStreamingSession: @unchecked Sendable {
audioArray: snapshot,
decodeOptions: optionsBuilder()
)
let text = results
.compactMap(\.text)
.joined(separator: " ")
.trimmingCharacters(in: .whitespacesAndNewlines)
let text = TranscriptSegmentJoiner.joined(
results.compactMap(\.text),
language: self?.language
)
await self?.finishPartialTask(text: text, submittedSampleCount: submittedSampleCount)
return text
}
Expand Down
15 changes: 15 additions & 0 deletions Tests/OpenTypeTests/SpeechRecognitionQualityTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,21 @@ final class SpeechRecognitionQualityTests: XCTestCase {
XCTAssertEqual(prompt, "Dictation terms: OpenType, MLX.")
}

func testWhisperPromptTokensInjectTermsWhenLanguageIsAuto() {
let context = SpeechRecognitionContext(phrases: ["OpenType", "MLX"])

let tokens = context.whisperPromptTokens(
language: nil,
maximumCount: 160,
tokenize: { Array($0.utf8).map(Int.init) }
)
let prompt = tokens.map { String(decoding: $0.map(UInt8.init), as: UTF8.self) }

XCTAssertNotNil(tokens)
XCTAssertFalse(tokens?.isEmpty ?? true)
XCTAssertEqual(prompt, "Dictation terms: OpenType, MLX.")
}

func testAppleCompatibleDictationPresetChangesAfterOneMinute() {
XCTAssertEqual(
AppleSpeechAnalyzer.dictationPreset(forDuration: 30),
Expand Down
46 changes: 46 additions & 0 deletions Tests/OpenTypeTests/TranscriptSegmentJoinerTests.swift
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
import XCTest
@testable import OpenType

final class TranscriptSegmentJoinerTests: XCTestCase {
func testCJKLanguageJoinsChunkSegmentsWithoutSpaces() {
XCTAssertEqual(
TranscriptSegmentJoiner.joined(["你好", "世界"], language: "zh"),
"你好世界"
)
XCTAssertEqual(
TranscriptSegmentJoiner.joined(["こんにちは", "世界"], language: "ja"),
"こんにちは世界"
)
XCTAssertEqual(
TranscriptSegmentJoiner.joined(["안녕하세요", "세계"], language: "ko"),
"안녕하세요세계"
)
}

func testLatinLanguageKeepsSpaceBetweenChunkSegments() {
XCTAssertEqual(
TranscriptSegmentJoiner.joined(["hello", "world"], language: "en"),
"hello world"
)
}

func testAutoLanguageFallsBackToBoundaryCharacters() {
XCTAssertEqual(
TranscriptSegmentJoiner.joined(["你好", "世界"], language: nil),
"你好世界"
)
XCTAssertEqual(
TranscriptSegmentJoiner.joined(["hello", "world"], language: nil),
"hello world"
)
}

func testTrimsSegmentsAndDropsEmpties() {
XCTAssertEqual(
TranscriptSegmentJoiner.joined([" 你好 ", "", " 世界 "], language: "zh"),
"你好世界"
)
XCTAssertEqual(TranscriptSegmentJoiner.joined([], language: "zh"), "")
XCTAssertEqual(TranscriptSegmentJoiner.joined(["", " "], language: "en"), "")
}
}
Loading