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
Original file line number Diff line number Diff line change
Expand Up @@ -88,18 +88,24 @@ extension StreamingNemotronMultilingualAsrManager {
) {
guard partialPublicationSuppressionDepth == 0, !inBlankRescue else { return }
guard let callback = callback, let tokenizer = tokenizer else { return }
guard !isFinal, Self.blankRescueEnabled, Self.rescueRmsThreshold > 0 else {
guard Self.blankRescueEnabled, Self.rescueRmsThreshold > 0 else {
callback(tokenizer.decode(ids: accumulatedTokenIds).text)
return
}
let ids = Self.partialPublicationTokenIds(
liveIds: accumulatedTokenIds, liveTimings: accumulatedTokenTimings,
langTagTokenIds: config.langTagTokenIds,
openSpan: rescueSpanOpen
? (rescueSpanStartFrame, rescueSpanLastSpeechFrame, rescueSpanPreRollFrames, rescueSpanOverflowed)
: nil,
nextSpanStartFrame: rescueFrameCursor - rescuePreRollTail.count / ASRConstants.samplesPerEncoderFrame)
callback(tokenizer.decode(ids: ids).text)
let ids =
isFinal
? accumulatedTokenIds
: Self.partialPublicationTokenIds(
liveIds: accumulatedTokenIds, liveTimings: accumulatedTokenTimings,
langTagTokenIds: config.langTagTokenIds,
openSpan: rescueSpanOpen
? (rescueSpanStartFrame, rescueSpanLastSpeechFrame, rescueSpanPreRollFrames, rescueSpanOverflowed)
: nil,
nextSpanStartFrame: rescueFrameCursor - rescuePreRollTail.count / ASRConstants.samplesPerEncoderFrame)
let text = tokenizer.decode(ids: ids).text
guard text != lastDeliveredPartialText else { return }
lastDeliveredPartialText = text
callback(text)
}

/// `processChunk` plus blank-span bookkeeping. All streaming call sites
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,8 @@ import AVFoundation
import Foundation

/// Callback delivering cumulative transcript text; language tags use `detectedLanguage()`.
/// With blank rescue enabled, text is settled and only extends the previous text;
/// chunks with nothing newly settled may repeat it unchanged. An unresolved blank
/// With blank rescue enabled, text is settled, only extends the previous text,
/// and is delivered only when it changes. An unresolved blank
/// span can hold back later text for the 15 s span cap (15.04 s at window granularity)
/// plus the rest of that chunk in audio time; processing time is additional.
/// `getPartialTranscript()` may already include that text; `finish()` flushes the rest.
Expand Down Expand Up @@ -269,6 +269,7 @@ public actor StreamingNemotronMultilingualAsrManager {
// Callbacks
internal var partialCallback: NemotronMultilingualPartialCallback?
internal var partialPublicationSuppressionDepth: Int = 0
internal var lastDeliveredPartialText: String?

// Stats
internal var processedChunks: Int = 0
Expand Down Expand Up @@ -813,6 +814,7 @@ public actor StreamingNemotronMultilingualAsrManager {
/// Reset all states for a new transcription session.
/// Preserves the currently selected prompt id and ml configuration.
public func reset() async {
lastDeliveredPartialText = nil
StreamingAsrUtils.resetSharedState(
audioBuffer: &audioBuffer,
accumulatedTokenIds: &accumulatedTokenIds,
Expand Down Expand Up @@ -903,6 +905,7 @@ public actor StreamingNemotronMultilingualAsrManager {
}

internal func resetStates() throws {
lastDeliveredPartialText = nil
let cacheConfig = EncoderCacheManager.CacheConfig(
channelShape: config.cacheChannelShape,
timeShape: config.cacheTimeShape,
Expand Down Expand Up @@ -1057,6 +1060,7 @@ public actor StreamingNemotronMultilingualAsrManager {
lastFinishTokenTimings = accumulatedTokenTimings
accumulatedTokenIds.removeAll()
accumulatedTokenTimings.removeAll()
lastDeliveredPartialText = nil
// The emitted-token tail must follow the accumulated ids it mirrors.
vocabularyBias?.resetMatchState()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,116 @@ final class NemotronMultilingualPublicationTests: XCTestCase {
private let later = 5
private let boundary = 6

func testRepeatedSettledPublicationDeliversOnce() async throws {
let tokenizer = try makeTokenizer()
let manager = Manager()
let updates = OSAllocatedUnfairLock<[String]>(initialState: [])
await manager.setPartialCallback { text in updates.withLock { $0.append(text) } }
await manager.setPublicationState(
tokenizer: tokenizer, ids: [hello, period], timings: [timing(hello, 5), timing(period, 17)],
openSpanStart: 14)

await manager.publishCurrentState()
await manager.publishCurrentState()

XCTAssertEqual(updates.withLock { $0 }, ["Hallo"])
}

func testChangedSettledTextDeliversAnotherUpdate() async throws {
let tokenizer = try makeTokenizer()
let manager = Manager()
let updates = OSAllocatedUnfairLock<[String]>(initialState: [])
await manager.setPartialCallback { text in updates.withLock { $0.append(text) } }
await manager.setPublicationState(tokenizer: tokenizer, ids: [hello], timings: [timing(hello, 5)])
await manager.publishCurrentState()
await manager.setPublicationState(
tokenizer: tokenizer, ids: [hello, world], timings: [timing(hello, 5), timing(world, 14)])
await manager.publishCurrentState()
await manager.publishCurrentState()

XCTAssertEqual(updates.withLock { $0 }, ["Hallo", "Hallo Welt"])
}

func testEqualDecodedTextAfterNewTokensPreservesRescueDisabledBehavior() async throws {
let tokenizer = try makeTokenizer()
let manager = Manager()
let updates = OSAllocatedUnfairLock<[String]>(initialState: [])
await manager.setPartialCallback { text in updates.withLock { $0.append(text) } }
await manager.setPublicationState(tokenizer: tokenizer, ids: [hello], timings: [timing(hello, 5)])
await manager.publishCurrentState()
await manager.setPublicationState(
tokenizer: tokenizer, ids: [hello, langTag], timings: [timing(hello, 5)])
await manager.publishCurrentState()
await manager.setPublicationState(
tokenizer: tokenizer, ids: [hello, langTag, boundary], timings: [timing(hello, 5), timing(boundary, 6)])
await manager.publishCurrentState()

let rescueEnabled = Manager.blankRescueEnabled && Manager.rescueRmsThreshold > 0
XCTAssertEqual(updates.withLock { $0 }, rescueEnabled ? ["Hallo"] : ["Hallo", "Hallo", "Hallo"])
}

func testResetAllowsTheSameTextToDeliverAgain() async throws {
let tokenizer = try makeTokenizer()
let manager = Manager()
let updates = OSAllocatedUnfairLock<[String]>(initialState: [])
await manager.setPartialCallback { text in updates.withLock { $0.append(text) } }
await manager.setPublicationState(tokenizer: tokenizer, ids: [hello], timings: [timing(hello, 5)])
await manager.publishCurrentState()

await manager.reset()
await manager.setPublicationState(tokenizer: tokenizer, ids: [hello], timings: [timing(hello, 5)])
await manager.publishCurrentState()
await manager.publishCurrentState()

XCTAssertEqual(updates.withLock { $0 }, ["Hallo", "Hallo"])
}

func testResetStatesAllowsTheSameTextToDeliverAgain() async throws {
let tokenizer = try makeTokenizer()
let manager = Manager()
let updates = OSAllocatedUnfairLock<[String]>(initialState: [])
await manager.setPartialCallback { text in updates.withLock { $0.append(text) } }
await manager.setPublicationState(tokenizer: tokenizer, ids: [hello], timings: [timing(hello, 5)])
await manager.publishCurrentState()

try await manager.resetStates()
await manager.setPublicationState(tokenizer: tokenizer, ids: [hello], timings: [timing(hello, 5)])
await manager.publishCurrentState()
await manager.publishCurrentState()

XCTAssertEqual(updates.withLock { $0 }, ["Hallo", "Hallo"])
}

func testFinalFlushDeliversChangedTextOnlyOnce() async throws {
let tokenizer = try makeTokenizer()
let manager = Manager()
let updates = OSAllocatedUnfairLock<[String]>(initialState: [])
await manager.setPartialCallback { text in updates.withLock { $0.append(text) } }
await manager.setPublicationState(
tokenizer: tokenizer, ids: [hello, period], timings: [timing(hello, 5), timing(period, 17)],
openSpanStart: 14)
await manager.publishCurrentState()

try await manager.finalizeRescueSpanIfNeeded()
try await manager.finalizeRescueSpanIfNeeded()

XCTAssertEqual(updates.withLock { $0 }, ["Hallo", "Hallo."])
}

func testFinalFlushSkipsTextAlreadyDeliveredLive() async throws {
let tokenizer = try makeTokenizer()
let manager = Manager()
let updates = OSAllocatedUnfairLock<[String]>(initialState: [])
await manager.setPartialCallback { text in updates.withLock { $0.append(text) } }
await manager.setPublicationState(
tokenizer: tokenizer, ids: [hello, period], timings: [timing(hello, 5), timing(period, 17)])
await manager.publishCurrentState()

try await manager.finalizeRescueSpanIfNeeded()

XCTAssertEqual(updates.withLock { $0 }, ["Hallo."])
}

func testRescueInsertionBeforeAccumulatedPunctuationPublishesOnlyExtensions() throws {
let tokenizer = try makeTokenizer()
let ids = [langTag, hello, period, later]
Expand Down Expand Up @@ -198,7 +308,7 @@ final class NemotronMultilingualPublicationTests: XCTestCase {
await manager.publishCurrentState()
await manager.publishCurrentState()
XCTAssertEqual(originalUpdates.withLock { $0 }, [])
XCTAssertEqual(replacementUpdates.withLock { $0 }, ["Hallo.", "Hallo."])
XCTAssertEqual(replacementUpdates.withLock { $0 }, ["Hallo."])
}

func testNestedSuppressionKeepsTheCurrentCallbackUntilTheChunkSettles() async throws {
Expand Down Expand Up @@ -228,7 +338,7 @@ final class NemotronMultilingualPublicationTests: XCTestCase {
let depth = await manager.partialPublicationSuppressionDepth
XCTAssertEqual(depth, 0)
XCTAssertEqual(originalUpdates.withLock { $0 }, [])
XCTAssertEqual(replacementUpdates.withLock { $0 }, ["Hallo.", "Hallo."])
XCTAssertEqual(replacementUpdates.withLock { $0 }, ["Hallo."])
}

private func timing(_ id: Int, _ frame: Int) -> TokenTiming {
Expand Down
Loading