From 4f6a869433c121912a75ecce807b4cd526ed340b Mon Sep 17 00:00:00 2001 From: james-333i Date: Tue, 25 Aug 2026 13:44:06 -0700 Subject: [PATCH] Deliver streamed tokens live instead of after generation completes streamResponse consumed an inner AsyncThrowingStream whose builder ran the entire generation loop synchronously on the consuming task, so every snapshot buffered and arrived in one burst after generation finished. Yield snapshots directly from the generation loop on the streaming task, and check for task cancellation between tokens so an abandoned stream stops decoding promptly. --- .../Models/LlamaLanguageModel.swift | 217 ++++++++---------- 1 file changed, 94 insertions(+), 123 deletions(-) diff --git a/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift b/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift index 8b8bb0bf..27325167 100644 --- a/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift @@ -644,25 +644,20 @@ import Foundation assistantPrefill: runtimeOptions.assistantPrefill ) - do { - for try await tokenText in generateTextStream( - context: context, - model: model!, - prompt: fullPrompt, - maxTokens: maxTokens, - options: runtimeOptions - ) { - accumulatedText += tokenText - - let snapshot = LanguageModelSession.ResponseStream.Snapshot( - content: (accumulatedText as! Content).asPartiallyGenerated(), - rawContent: GeneratedContent(accumulatedText) - ) - continuation.yield(snapshot) - } - } catch { - continuation.finish(throwing: error) - return + try self.performTextGeneration( + context: context, + model: model!, + prompt: fullPrompt, + maxTokens: maxTokens, + options: runtimeOptions + ) { tokenText in + accumulatedText += tokenText + + let snapshot = LanguageModelSession.ResponseStream.Snapshot( + content: (accumulatedText as! Content).asPartiallyGenerated(), + rawContent: GeneratedContent(accumulatedText) + ) + continuation.yield(snapshot) } continuation.finish() @@ -1186,133 +1181,109 @@ import Foundation } } - private func generateTextStream( - context: OpaquePointer, - model: OpaquePointer, - prompt: String, - maxTokens: Int, - options: ResolvedGenerationOptions - ) -> AsyncThrowingStream { - return AsyncThrowingStream { continuation in - self.performTextGeneration( - context: context, - model: model, - prompt: prompt, - maxTokens: maxTokens, - options: options, - continuation: continuation - ) - } - } - private func performTextGeneration( context: OpaquePointer, model: OpaquePointer, prompt: String, maxTokens: Int, options: ResolvedGenerationOptions, - continuation: AsyncThrowingStream.Continuation - ) { - do { - guard let vocab = llama_model_get_vocab(model) else { - continuation.finish(throwing: LlamaLanguageModelError.contextInitializationFailed) - return - } + onToken: (String) -> Void + ) throws { + guard let vocab = llama_model_get_vocab(model) else { + throw LlamaLanguageModelError.contextInitializationFailed + } - // Tokenize the prompt - let promptTokens = try tokenizeText(vocab: vocab, text: prompt) - guard !promptTokens.isEmpty else { - continuation.finish(throwing: LlamaLanguageModelError.tokenizationFailed) - return - } + // Tokenize the prompt + let promptTokens = try tokenizeText(vocab: vocab, text: prompt) + guard !promptTokens.isEmpty else { + throw LlamaLanguageModelError.tokenizationFailed + } - // Initialize batch - var batch = llama_batch_init(Int32(options.batchSize), 0, 1) - defer { llama_batch_free(batch) } + // Initialize batch + var batch = llama_batch_init(Int32(options.batchSize), 0, 1) + defer { llama_batch_free(batch) } - let hasEncoder = try prepareInitialBatch( - batch: &batch, - promptTokens: promptTokens, - model: model, - vocab: vocab, - context: context, - batchSize: options.batchSize, - contextSize: options.contextSize - ) + let hasEncoder = try prepareInitialBatch( + batch: &batch, + promptTokens: promptTokens, + model: model, + vocab: vocab, + context: context, + batchSize: options.batchSize, + contextSize: options.contextSize + ) - // Initialize sampler chain with options - guard let sampler = llama_sampler_chain_init(llama_sampler_chain_default_params()) else { - throw LlamaLanguageModelError.decodingFailed - } - defer { llama_sampler_free(sampler) } - let samplerPtr = UnsafeMutablePointer(sampler) + // Initialize sampler chain with options + guard let sampler = llama_sampler_chain_init(llama_sampler_chain_default_params()) else { + throw LlamaLanguageModelError.decodingFailed + } + defer { llama_sampler_free(sampler) } + let samplerPtr = UnsafeMutablePointer(sampler) - let effectiveTemperature = Float(options.temperature) + let effectiveTemperature = Float(options.temperature) - // Apply repeat/frequency/presence penalties from custom options - let effectiveRepeatPenalty = options.repeatPenalty - let effectiveRepeatLastN = options.repeatLastN - let effectiveFrequencyPenalty = options.frequencyPenalty - let effectivePresencePenalty = options.presencePenalty + // Apply repeat/frequency/presence penalties from custom options + let effectiveRepeatPenalty = options.repeatPenalty + let effectiveRepeatLastN = options.repeatLastN + let effectiveFrequencyPenalty = options.frequencyPenalty + let effectivePresencePenalty = options.presencePenalty - if effectiveRepeatPenalty != 1.0 || effectiveFrequencyPenalty != 0.0 || effectivePresencePenalty != 0.0 - { - llama_sampler_chain_add( - samplerPtr, - llama_sampler_init_penalties( - llama_vocab_n_tokens(vocab), - effectiveRepeatLastN, - effectiveRepeatPenalty, - effectiveFrequencyPenalty, - effectivePresencePenalty - ) + if effectiveRepeatPenalty != 1.0 || effectiveFrequencyPenalty != 0.0 || effectivePresencePenalty != 0.0 { + llama_sampler_chain_add( + samplerPtr, + llama_sampler_init_penalties( + llama_vocab_n_tokens(vocab), + effectiveRepeatLastN, + effectiveRepeatPenalty, + effectiveFrequencyPenalty, + effectivePresencePenalty ) - } - - // Check for mirostat sampling (takes precedence over standard sampling) - applySampling(sampler: samplerPtr, effectiveTemperature: effectiveTemperature, options: options) + ) + } - // Generate tokens one by one - // Track position - for encoder-decoder models, we start from position 1 (after decoder start token) - // For decoder-only models, we continue from the end of the prompt - var n_cur: Int32 = hasEncoder ? 1 : Int32(promptTokens.count) + // Check for mirostat sampling (takes precedence over standard sampling) + applySampling(sampler: samplerPtr, effectiveTemperature: effectiveTemperature, options: options) - for _ in 0 ..< maxTokens { - // Sample next token from logits of the last token we just decoded - let nextToken = llama_sampler_sample(sampler, context, batch.n_tokens - 1) - llama_sampler_accept(sampler, nextToken) + // Generate tokens one by one + // Track position - for encoder-decoder models, we start from position 1 (after decoder start token) + // For decoder-only models, we continue from the end of the prompt + var n_cur: Int32 = hasEncoder ? 1 : Int32(promptTokens.count) - // Check for end of sequence - if llama_vocab_is_eog(vocab, nextToken) { - break - } + for _ in 0 ..< maxTokens { + if Task.isCancelled { + break + } - // Convert token to text and yield it - if let tokenText = tokenToText(vocab: vocab, token: nextToken) { - continuation.yield(tokenText) - } + // Sample next token from logits of the last token we just decoded + let nextToken = llama_sampler_sample(sampler, context, batch.n_tokens - 1) + llama_sampler_accept(sampler, nextToken) - // Prepare batch for next token - batch.n_tokens = 1 - batch.token[0] = nextToken - batch.pos[0] = n_cur - batch.n_seq_id[0] = 1 - if let seq_ids = batch.seq_id, let seq_id = seq_ids[0] { - seq_id[0] = 0 - } - batch.logits[0] = 1 + // Check for end of sequence + if llama_vocab_is_eog(vocab, nextToken) { + break + } - n_cur += 1 + // Convert token to text and yield it + if let tokenText = tokenToText(vocab: vocab, token: nextToken) { + onToken(tokenText) + } - let decodeResult = llama_decode(context, batch) - guard decodeResult == 0 else { - break - } + // Prepare batch for next token + batch.n_tokens = 1 + batch.token[0] = nextToken + batch.pos[0] = n_cur + batch.n_seq_id[0] = 1 + if let seq_ids = batch.seq_id, let seq_id = seq_ids[0] { + seq_id[0] = 0 } + batch.logits[0] = 1 - continuation.finish() - } catch { - continuation.finish(throwing: error) + n_cur += 1 + + let decodeResult = llama_decode(context, batch) + guard decodeResult == 0 else { + break + } } }