Skip to content
Open
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
133 changes: 84 additions & 49 deletions Sources/FluidAudio/Diarizer/Offline/Core/OfflineDiarizerManager.swift
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,12 @@ import Accelerate
import Foundation
import OSLog

@available(macOS 14.0, iOS 17.0, *)
private enum OfflinePreparationWorkerResult: Sendable {
case segmentation(SegmentationOutput, TimeInterval)
case embeddings([TimedEmbedding], TimeInterval)
}

@available(macOS 14.0, iOS 17.0, *)
public final class OfflineDiarizerManager {
private let logger = AppLogger(category: "OfflineDiarizer")
Expand Down Expand Up @@ -187,64 +193,93 @@ public final class OfflineDiarizerManager {
let capturedModels = models
let capturedConfig = config

let segmentationTask = Task.detached(priority: .userInitiated) {
[capturedModels, capturedConfig] () throws -> (SegmentationOutput, TimeInterval) in
let processor = OfflineSegmentationProcessor()
let start = Date()
do {
let segmentation = try await processor.process(
audioSource: audioSource,
segmentationModel: capturedModels.segmentationModel,
config: capturedConfig,
chunkHandler: { chunk in
progressCallback?(chunk.chunkIndex + 1, totalChunks)
switch chunkContinuation.yield(chunk) {
case .enqueued, .dropped:
return .continue
case .terminated:
return .stop
@unknown default:
return .stop
let results = try await withThrowingTaskGroup(
of: OfflinePreparationWorkerResult.self,
returning: (
segmentation: (SegmentationOutput, TimeInterval),
embeddings: ([TimedEmbedding], TimeInterval)
).self
) { group in
group.addTask(priority: .userInitiated) { [capturedModels, capturedConfig] in
let processor = OfflineSegmentationProcessor()
let start = Date()
do {
let segmentation = try await processor.process(
audioSource: audioSource,
segmentationModel: capturedModels.segmentationModel,
config: capturedConfig,
chunkHandler: { chunk in
progressCallback?(chunk.chunkIndex + 1, totalChunks)
switch chunkContinuation.yield(chunk) {
case .enqueued, .dropped:
return .continue
case .terminated:
return .stop
@unknown default:
return .stop
}
}
}
)
chunkContinuation.finish()
return .segmentation(
segmentation,
Date().timeIntervalSince(start)
)
} catch {
chunkContinuation.finish(throwing: error)
throw error
}
}

group.addTask(priority: .userInitiated) { [capturedModels, capturedConfig] in
let extractor = OfflineEmbeddingExtractor(
fbankModel: capturedModels.fbankModel,
embeddingModel: capturedModels.embeddingModel,
pldaTransform: PLDATransform(
pldaRhoModel: capturedModels.pldaRhoModel,
psi: capturedModels.pldaPsi
),
config: capturedConfig
)
let start = Date()
let embeddings = try await extractor.extractEmbeddings(
audioSource: audioSource,
segmentationStream: chunkStream
)
return .embeddings(
embeddings,
Date().timeIntervalSince(start)
)
chunkContinuation.finish()
return (segmentation, Date().timeIntervalSince(start))
}

var segmentationResult: (SegmentationOutput, TimeInterval)?
var embeddingResult: ([TimedEmbedding], TimeInterval)?

do {
while let result = try await group.next() {
switch result {
case .segmentation(let segmentation, let duration):
segmentationResult = (segmentation, duration)
case .embeddings(let embeddings, let duration):
embeddingResult = (embeddings, duration)
}
}
} catch {
group.cancelAll()
chunkContinuation.finish(throwing: error)
throw error
}
}

let embeddingTask = Task.detached(priority: .userInitiated) {
[capturedModels, capturedConfig] () throws -> ([TimedEmbedding], TimeInterval) in
let extractor = OfflineEmbeddingExtractor(
fbankModel: capturedModels.fbankModel,
embeddingModel: capturedModels.embeddingModel,
pldaTransform: PLDATransform(pldaRhoModel: capturedModels.pldaRhoModel, psi: capturedModels.pldaPsi),
config: capturedConfig
)
let start = Date()
let embeddings = try await extractor.extractEmbeddings(
audioSource: audioSource,
segmentationStream: chunkStream
)
return (embeddings, Date().timeIntervalSince(start))
guard let segmentationResult, let embeddingResult else {
throw OfflineDiarizationError.processingFailed(
"Offline preparation workers ended without complete results"
)
}
return (segmentationResult, embeddingResult)
}

let segmentationResult: (SegmentationOutput, TimeInterval)
let embeddingResult: ([TimedEmbedding], TimeInterval)
do {
async let awaitedSegmentation = segmentationTask.value
async let awaitedEmbeddings = embeddingTask.value
segmentationResult = try await awaitedSegmentation
embeddingResult = try await awaitedEmbeddings
} catch {
segmentationTask.cancel()
embeddingTask.cancel()
chunkContinuation.finish(throwing: error)
throw error
}
let segmentationResult = results.segmentation
let embeddingResult = results.embeddings

let (segmentation, segmentationTime) = segmentationResult
logger.debug("Segmentation completed in \(segmentationTime)s (async)")
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
import os
import XCTest

@testable import FluidAudio

@available(macOS 14.0, iOS 17.0, *)
final class OfflineDiarizerCancellationTests: XCTestCase {
func testCancellingFileProcessingStopsAllInferenceWorkers() async throws {
try requireOfflineDiarizerModels()

guard
let audioPath = ProcessInfo.processInfo.environment[
"FLUIDAUDIO_CANCELLATION_TEST_AUDIO"
],
FileManager.default.fileExists(atPath: audioPath)
else {
throw XCTSkip(
"Set FLUIDAUDIO_CANCELLATION_TEST_AUDIO to a real meeting recording"
)
}
let audioURL = URL(fileURLWithPath: audioPath)

let manager = OfflineDiarizerManager()
let progress = AsyncStream<Void>.makeStream(bufferingPolicy: .bufferingNewest(1))
let progressCount = OSAllocatedUnfairLock<Int>(initialState: 0)
var progressIterator = progress.stream.makeAsyncIterator()

let processingTask = Task {
try await manager.process(audioURL) { _, _ in
progressCount.withLock { $0 += 1 }
progress.continuation.yield(())
}
}

guard await progressIterator.next() != nil else {
processingTask.cancel()
XCTFail("Diarization ended before inference began")
return
}

let cancellationStart = ContinuousClock.now
processingTask.cancel()

do {
_ = try await processingTask.value
XCTFail("Cancelled diarization unexpectedly completed")
} catch is CancellationError {
// Expected. Returning from the structured task group also proves both
// inference workers have stopped rather than continuing in the background.
} catch {
XCTFail("Expected CancellationError, received \(error)")
}

let cancellationDuration = cancellationStart.duration(to: .now)
XCTAssertLessThan(
cancellationDuration,
.seconds(2),
"Cancellation should stop segmentation and embedding between model calls"
)

let countWhenCancelledCallReturned = progressCount.withLock { $0 }
try await Task.sleep(for: .milliseconds(250))
XCTAssertEqual(
progressCount.withLock { $0 },
countWhenCancelledCallReturned,
"No inference progress may continue after process(URL) returns"
)
}

private func requireOfflineDiarizerModels() throws {
let repoDirectory = OfflineDiarizerModels.defaultModelsDirectory()
.appendingPathComponent(Repo.diarizer.folderName, isDirectory: true)
let allPresent = ModelNames.OfflineDiarizer.requiredModels.allSatisfy {
FileManager.default.fileExists(atPath: repoDirectory.appendingPathComponent($0).path)
}
guard allPresent else {
throw XCTSkip("Offline diarizer models not available")
}
}
}
Loading