diff --git a/Package.resolved b/Package.resolved index b689a691..0b86bc46 100644 --- a/Package.resolved +++ b/Package.resolved @@ -19,6 +19,15 @@ "version" : "1.3.1" } }, + { + "identity" : "llama.swift", + "kind" : "remoteSourceControl", + "location" : "https://github.com/mattt/llama.swift", + "state" : { + "revision" : "716419d4d7aa542fce301e809cde7234c68ddbc6", + "version" : "2.10549.0" + } + }, { "identity" : "partialjsondecoder", "kind" : "remoteSourceControl", diff --git a/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift b/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift index 6bcd8045..e20ca8d1 100644 --- a/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift @@ -137,6 +137,14 @@ import Foundation /// Mirostat sampling mode for adaptive perplexity control. public var mirostat: MirostatMode? + /// Text appended after the assistant header of the rendered prompt. + /// + /// The model continues generating from this text. Use it to steer the + /// start of the response, for example prefilling an empty + /// `` block to suppress a model's default reasoning + /// output when its chat template offers no switch for it. + public var assistantPrefill: String? + /// Creates custom generation options for llama.cpp. public init( contextSize: UInt32? = nil, @@ -150,7 +158,8 @@ import Foundation repeatLastN: Int32? = nil, frequencyPenalty: Float? = nil, presencePenalty: Float? = nil, - mirostat: MirostatMode? = nil + mirostat: MirostatMode? = nil, + assistantPrefill: String? = nil ) { self.contextSize = contextSize self.batchSize = batchSize @@ -164,6 +173,7 @@ import Foundation self.frequencyPenalty = frequencyPenalty self.presencePenalty = presencePenalty self.mirostat = mirostat + self.assistantPrefill = assistantPrefill } /// Default llama.cpp options used when none are provided at runtime. @@ -321,6 +331,7 @@ import Foundation var frequencyPenalty: Float var presencePenalty: Float var mirostat: CustomGenerationOptions.MirostatMode? + var assistantPrefill: String? var sampling: GenerationOptions.SamplingMode? var maximumResponseTokens: Int? @@ -337,6 +348,7 @@ import Foundation frequencyPenalty: Float = 0.0, presencePenalty: Float = 0.0, mirostat: CustomGenerationOptions.MirostatMode? = nil, + assistantPrefill: String? = nil, sampling: GenerationOptions.SamplingMode? = nil, maximumResponseTokens: Int? = nil ) { @@ -352,6 +364,7 @@ import Foundation self.frequencyPenalty = frequencyPenalty self.presencePenalty = presencePenalty self.mirostat = mirostat + self.assistantPrefill = assistantPrefill self.sampling = sampling self.maximumResponseTokens = maximumResponseTokens } @@ -389,6 +402,7 @@ import Foundation frequencyPenalty: base.frequencyPenalty, presencePenalty: base.presencePenalty, mirostat: base.mirostat, + assistantPrefill: base.assistantPrefill, sampling: sampling ?? base.sampling, maximumResponseTokens: maximumResponseTokens ?? base.maximumResponseTokens ) @@ -407,6 +421,7 @@ import Foundation self.frequencyPenalty = options.frequencyPenalty ?? base.frequencyPenalty self.presencePenalty = options.presencePenalty ?? base.presencePenalty self.mirostat = options.mirostat ?? base.mirostat + self.assistantPrefill = options.assistantPrefill ?? base.assistantPrefill self.sampling = sampling ?? base.sampling self.maximumResponseTokens = maximumResponseTokens ?? base.maximumResponseTokens } @@ -509,10 +524,11 @@ import Foundation if includeSchemaInPrompt, type != String.self { fullPrompt = try formatPrompt( for: session, - extraSystemMessage: schemaPrompt(for: type.generationSchema) + extraSystemMessage: schemaPrompt(for: type.generationSchema), + assistantPrefill: runtimeOptions.assistantPrefill ) } else { - fullPrompt = try formatPrompt(for: session) + fullPrompt = try formatPrompt(for: session, assistantPrefill: runtimeOptions.assistantPrefill) } if type == String.self { @@ -600,7 +616,10 @@ import Foundation llama_set_n_threads(context, runtimeOptions.threads, runtimeOptions.threads) var accumulatedText = "" - let fullPrompt = try self.formatPrompt(for: session) + let fullPrompt = try self.formatPrompt( + for: session, + assistantPrefill: runtimeOptions.assistantPrefill + ) do { for try await tokenText in generateTextStream( @@ -673,8 +692,7 @@ import Foundation params.n_gpu_layers = 0 // Try to reduce memory usage - params.use_mmap = true - params.use_mlock = false + params.load_mode = LLAMA_LOAD_MODE_MMAP return params } @@ -829,6 +847,7 @@ import Foundation llama_sampler_chain_add( samplerPtr, llama_sampler_init_penalties( + llama_vocab_n_tokens(vocab), effectiveRepeatLastN, effectiveRepeatPenalty, effectiveFrequencyPenalty, @@ -958,6 +977,7 @@ import Foundation llama_sampler_chain_add( samplerPointer, llama_sampler_init_penalties( + llama_vocab_n_tokens(vocab), options.repeatLastN, options.repeatPenalty, options.frequencyPenalty, @@ -1197,6 +1217,7 @@ import Foundation llama_sampler_chain_add( samplerPtr, llama_sampler_init_penalties( + llama_vocab_n_tokens(vocab), effectiveRepeatLastN, effectiveRepeatPenalty, effectiveFrequencyPenalty, @@ -1368,7 +1389,8 @@ import Foundation private func formatPrompt( for session: LanguageModelSession, - extraSystemMessage: String? = nil + extraSystemMessage: String? = nil, + assistantPrefill: String? = nil ) throws -> String { guard let model = self.model else { throw LlamaLanguageModelError.modelLoadFailed @@ -1452,9 +1474,14 @@ import Foundation throw LlamaLanguageModelError.encodingFailed } - return buffer.withUnsafeBytes { rawBuffer in + let rendered = buffer.withUnsafeBytes { rawBuffer in String(decoding: rawBuffer.prefix(Int(result)), as: UTF8.self) } + + if let assistantPrefill, !assistantPrefill.isEmpty { + return rendered + assistantPrefill + } + return rendered } private func extractText(from segments: [Transcript.Segment]) -> String { diff --git a/Tests/AnyLanguageModelTests/LlamaLanguageModelTests.swift b/Tests/AnyLanguageModelTests/LlamaLanguageModelTests.swift index 3d9e32c7..fa2ae7c2 100644 --- a/Tests/AnyLanguageModelTests/LlamaLanguageModelTests.swift +++ b/Tests/AnyLanguageModelTests/LlamaLanguageModelTests.swift @@ -45,7 +45,8 @@ import Testing repeatLastN: 48, frequencyPenalty: 0.05, presencePenalty: 0.05, - mirostat: .v2(tau: 5.0, eta: 0.2) + mirostat: .v2(tau: 5.0, eta: 0.2), + assistantPrefill: "" ) options[custom: LlamaLanguageModel.self] = custom @@ -62,6 +63,7 @@ import Testing #expect(retrieved?.frequencyPenalty == 0.05) #expect(retrieved?.presencePenalty == 0.05) #expect(retrieved?.mirostat == .v2(tau: 5.0, eta: 0.2)) + #expect(retrieved?.assistantPrefill == "") } @Test func customGenerationOptionsDefaults() {