Skip to content

Commit 4d54ff3

Browse files
committed
Make ANE utils concurrency safe
1 parent 5bb7cdd commit 4d54ff3

9 files changed

Lines changed: 243 additions & 99 deletions

File tree

Documentation/API.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@ This page summarizes the primary public APIs across modules. See inline doc comm
1818
Main class for speaker diarization and "who spoke when" analysis.
1919

2020
**Key Methods:**
21-
- `performCompleteDiarization(_:sampleRate:) throws -> DiarizerResult`
21+
- `performCompleteDiarization(_:sampleRate:) throws -> DiarizationResult`
2222
- Process complete audio file and return speaker segments
2323
- Parameters: `RandomAccessCollection<Float>` audio samples, sample rate (default: 16000)
2424
- Returns: `DiarizerResult` with speaker segments and timing

Sources/FluidAudio/Diarizer/Core/DiarizerManager.swift

Lines changed: 45 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import Accelerate
22
import CoreML
3+
import Dispatch
34
import Foundation
45
import OSLog
56

@@ -8,8 +9,6 @@ public final class DiarizerManager {
89
internal let logger = AppLogger(category: "Diarizer")
910
internal let config: DiarizerConfig
1011
private var models: DiarizerModels?
11-
private var chunkBuffer: [Float] = []
12-
1312
/// Public getter for segmentation model (for streaming)
1413
public var segmentationModel: MLModel? {
1514
return models?.segmentationModel
@@ -126,6 +125,7 @@ public final class DiarizerManager {
126125
let overlapDuration = Int(config.chunkOverlap.rounded())
127126
let chunkSize = sampleRate * chunkDuration
128127
let stepSize = chunkSize - (sampleRate * overlapDuration)
128+
var chunkBuffer = [Float](repeating: 0.0, count: max(chunkSize, 1))
129129

130130
var allSegments: [TimedSpeakerSegment] = []
131131

@@ -145,7 +145,9 @@ public final class DiarizerManager {
145145
chunk,
146146
chunkOffset: chunkOffset,
147147
models: models,
148-
sampleRate: sampleRate
148+
sampleRate: sampleRate,
149+
chunkSize: chunkSize,
150+
chunkBuffer: &chunkBuffer
149151
)
150152
allSegments.append(contentsOf: chunkSegments)
151153

@@ -203,12 +205,13 @@ public final class DiarizerManager {
203205
_ chunk: C,
204206
chunkOffset: Double,
205207
models: DiarizerModels,
206-
sampleRate: Int = 16000
208+
sampleRate: Int = 16000,
209+
chunkSize: Int,
210+
chunkBuffer: inout [Float]
207211
) throws -> ([TimedSpeakerSegment], ChunkTimings)
208212
where C: RandomAccessCollection, C.Element == Float, C.Index == Int {
209213
let segmentationStartTime = Date()
210214

211-
let chunkSize = sampleRate * 10
212215
let chunkCount = chunk.distance(from: chunk.startIndex, to: chunk.endIndex)
213216
let copyCount = min(chunkCount, chunkSize)
214217

@@ -278,8 +281,9 @@ public final class DiarizerManager {
278281
masks.append(speakerMask)
279282
}
280283

281-
let embeddings = try embeddingExtractor.getEmbeddings(
282-
audio: Array(paddedChunk),
284+
let embeddings = try extractEmbeddingsSynchronously(
285+
using: embeddingExtractor,
286+
audio: paddedChunk,
283287
masks: masks,
284288
minActivityThreshold: config.minActiveFramesCount
285289
)
@@ -342,6 +346,40 @@ public final class DiarizerManager {
342346
return (segments, timings)
343347
}
344348

349+
private func extractEmbeddingsSynchronously<C>(
350+
using extractor: EmbeddingExtractor,
351+
audio: C,
352+
masks: [[Float]],
353+
minActivityThreshold: Float
354+
) throws -> [[Float]]
355+
where C: RandomAccessCollection, C.Element == Float, C.Index == Int {
356+
var embeddingsResult: [[Float]]?
357+
var capturedError: Error?
358+
let semaphore = DispatchSemaphore(value: 0)
359+
360+
Task {
361+
do {
362+
let result = try await extractor.getEmbeddings(
363+
audio: audio,
364+
masks: masks,
365+
minActivityThreshold: minActivityThreshold
366+
)
367+
embeddingsResult = result
368+
} catch {
369+
capturedError = error
370+
}
371+
semaphore.signal()
372+
}
373+
374+
semaphore.wait()
375+
376+
if let error = capturedError {
377+
throw error
378+
}
379+
380+
return embeddingsResult ?? []
381+
}
382+
345383
/// Count activity frames per speaker.
346384
private func calculateSpeakerActivities(_ binarizedSegments: [[[Float]]]) -> [Float] {
347385
let numSpeakers = binarizedSegments[0][0].count

Sources/FluidAudio/Diarizer/Extraction/EmbeddingExtractor.swift

Lines changed: 31 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -3,38 +3,17 @@ import CoreML
33
import OSLog
44

55
/// Embedding extractor with ANE-aligned memory and zero-copy operations
6-
public class EmbeddingExtractor {
6+
public actor EmbeddingExtractor {
77
private let wespeakerModel: MLModel
88
private let logger = AppLogger(category: "EmbeddingExtractor")
99
private let memoryOptimizer = ANEMemoryOptimizer()
1010

11-
// Pre-allocated ANE-aligned buffers
12-
private var waveformBuffer: MLMultiArray?
13-
private var maskBuffer: MLMultiArray?
14-
15-
// Reusable feature providers
16-
private var featureProviders: [ZeroCopyDiarizerFeatureProvider] = []
11+
private let waveformShape: [NSNumber] = [3, 160_000]
12+
private let waveformKey = "wespeaker_waveform_primary"
1713

1814
public init(embeddingModel: MLModel) {
1915
self.wespeakerModel = embeddingModel
20-
21-
// Pre-allocate ANE-aligned buffers
22-
do {
23-
self.waveformBuffer = try memoryOptimizer.createAlignedArray(
24-
shape: [3, 160000] as [NSNumber],
25-
dataType: .float32
26-
)
27-
28-
self.maskBuffer = try memoryOptimizer.createAlignedArray(
29-
shape: [3, 1000] as [NSNumber], // Typical mask size
30-
dataType: .float32
31-
)
32-
} catch {
33-
logger.error("Failed to allocate ANE-aligned buffers: \(error)")
34-
// Buffers will remain nil, will be allocated on-demand in getEmbeddings
35-
}
36-
37-
logger.info("EmbeddingExtractor initialized with ANE-aligned buffers")
16+
logger.info("EmbeddingExtractor ready with ANE memory optimizer")
3817
}
3918

4019
/// Extract speaker embeddings using the CoreML embedding model.
@@ -54,18 +33,32 @@ public class EmbeddingExtractor {
5433
minActivityThreshold: Float = 10.0
5534
) throws -> [[Float]]
5635
where C: RandomAccessCollection, C.Element == Float, C.Index == Int {
57-
// We need to return embeddings for ALL speakers, not just active ones
58-
// to maintain compatibility with the rest of the pipeline
59-
var embeddings: [[Float]] = []
36+
guard let firstMask = masks.first else {
37+
return []
38+
}
39+
40+
let maskShape = [3, firstMask.count] as [NSNumber]
41+
let maskKey = "wespeaker_mask_\(firstMask.count)"
42+
43+
let waveformLease = try memoryOptimizer.leaseBuffer(
44+
key: waveformKey,
45+
shape: waveformShape,
46+
dataType: .float32
47+
)
6048

61-
// Get or create appropriately sized mask buffer
62-
let maskShape = [3, masks[0].count] as [NSNumber]
63-
let currentMaskBuffer = try memoryOptimizer.getPooledBuffer(
64-
key: "wespeaker_mask_\(masks[0].count)",
49+
let maskLease = try memoryOptimizer.leaseBuffer(
50+
key: maskKey,
6551
shape: maskShape,
6652
dataType: .float32
6753
)
6854

55+
let waveformBuffer = waveformLease.multiArray
56+
let maskBuffer = maskLease.multiArray
57+
58+
// We need to return embeddings for ALL speakers, not just active ones
59+
// to maintain compatibility with the rest of the pipeline
60+
var embeddings: [[Float]] = []
61+
6962
// Process all speakers but optimize for active ones
7063
for speakerIdx in 0..<masks.count {
7164
// Check if speaker is active
@@ -80,28 +73,28 @@ public class EmbeddingExtractor {
8073
// Use ANE-optimized copy for audio data
8174
memoryOptimizer.optimizedCopy(
8275
from: audio,
83-
to: waveformBuffer!,
76+
to: waveformBuffer,
8477
offset: 0 // First speaker slot
8578
)
8679

8780
// Optimize mask creation with zero-copy view
8881
fillMaskBufferOptimized(
8982
masks: masks,
9083
speakerIndex: speakerIdx,
91-
buffer: currentMaskBuffer
84+
buffer: maskBuffer
9285
)
9386

9487
// Create zero-copy feature provider
9588
let featureProvider = ZeroCopyDiarizerFeatureProvider(features: [
96-
"waveform": MLFeatureValue(multiArray: waveformBuffer!),
97-
"mask": MLFeatureValue(multiArray: currentMaskBuffer),
89+
"waveform": MLFeatureValue(multiArray: waveformBuffer),
90+
"mask": MLFeatureValue(multiArray: maskBuffer),
9891
])
9992

10093
// Run model with optimal prediction options
10194
let options = MLPredictionOptions()
10295
// Prefetch to Neural Engine for better performance
103-
waveformBuffer!.prefetchToNeuralEngine()
104-
currentMaskBuffer.prefetchToNeuralEngine()
96+
waveformBuffer.prefetchToNeuralEngine()
97+
maskBuffer.prefetchToNeuralEngine()
10598

10699
let output = try wespeakerModel.prediction(from: featureProvider, options: options)
107100

Sources/FluidAudio/Diarizer/Offline/Extraction/OfflineEmbeddingExtractor.swift

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -259,20 +259,21 @@ struct OfflineEmbeddingExtractor {
259259
var pldaBatchCallCount = 0
260260

261261
func performEmbeddingWarmup() throws {
262-
let warmupAudioArray = try memoryOptimizer.getPooledBuffer(
262+
let warmupAudioLease = try memoryOptimizer.leaseBuffer(
263263
key: "offline_embedding_warmup_audio",
264264
shape: fbankInputShape,
265265
dataType: .float32
266266
)
267+
let warmupAudioArray = warmupAudioLease.multiArray
267268
let warmupAudioPointer = warmupAudioArray.dataPointer.assumingMemoryBound(to: Float.self)
268269
vDSP_vclr(warmupAudioPointer, 1, vDSP_Length(warmupAudioArray.count))
269270

270271
let warmupFbankFeatures = try runFbankModel(audioArray: warmupAudioArray)
271272
let zeroWeights = [Float](repeating: 0, count: weightFrameCount)
272-
let warmupWeightsArray = try prepareWeightsInput(weights: zeroWeights)
273+
let warmupWeightsLease = try prepareWeightsInput(weights: zeroWeights)
273274
_ = try runEmbeddingModel(
274275
fbankFeatures: warmupFbankFeatures,
275-
weightsArray: warmupWeightsArray
276+
weightsArray: warmupWeightsLease.multiArray
276277
)
277278
}
278279

@@ -455,10 +456,10 @@ struct OfflineEmbeddingExtractor {
455456
}
456457

457458
let embeddingStart = resampleEnd
458-
let weightsArray = try prepareWeightsInput(weights: resampledMask)
459+
let weightsLease = try prepareWeightsInput(weights: resampledMask)
459460
let embedding256 = try runEmbeddingModel(
460461
fbankFeatures: fbankFeatures,
461-
weightsArray: weightsArray
462+
weightsArray: weightsLease.multiArray
462463
)
463464
let embeddingEnd = clock.now
464465
embeddingDuration += embeddingStart.duration(to: embeddingEnd)
@@ -702,12 +703,13 @@ struct OfflineEmbeddingExtractor {
702703

703704
private func prepareWeightsInput(
704705
weights: [Float]
705-
) throws -> MLMultiArray {
706-
let array = try memoryOptimizer.getPooledBuffer(
706+
) throws -> ANEMultiArrayLease {
707+
let lease = try memoryOptimizer.leaseBuffer(
707708
key: "offline_embedding_weights_\(weightFrameCount)",
708709
shape: weightInputShape,
709710
dataType: .float32
710711
)
712+
let array = lease.multiArray
711713

712714
let pointer = array.dataPointer.assumingMemoryBound(to: Float.self)
713715
vDSP_vclr(pointer, 1, vDSP_Length(array.count))
@@ -726,7 +728,7 @@ struct OfflineEmbeddingExtractor {
726728
}
727729
}
728730

729-
return array
731+
return lease
730732
}
731733

732734
private func runEmbeddingModel(

0 commit comments

Comments
 (0)