diff --git a/LLama.Examples/NewVersion/LoadAndSaveState.cs b/LLama.Examples/NewVersion/LoadAndSaveState.cs index 7e7d033d..bf5e9daa 100644 --- a/LLama.Examples/NewVersion/LoadAndSaveState.cs +++ b/LLama.Examples/NewVersion/LoadAndSaveState.cs @@ -10,14 +10,14 @@ namespace LLama.Examples.NewVersion { public class SaveAndLoadState : IDisposable { - LLamaInteractExecutor _executor; + InteractiveExecutor _executor; string _prompt; string _modelPath; public SaveAndLoadState(string modelPath, string prompt) { _prompt = prompt; _modelPath = modelPath; - _executor = new LLamaInteractExecutor(new LLamaModel(new ModelParams(modelPath: modelPath))); + _executor = new InteractiveExecutor(new LLamaModel(new ModelParams(modelPath: modelPath))); foreach (var text in _executor.Infer(_prompt, new SessionParams() { Temperature = 0.6f, AntiPrompts = new List { "user:" } })) { Console.Write(text); @@ -44,7 +44,7 @@ namespace LLama.Examples.NewVersion { var model = _executor.Model; model.LoadState(modelStateFile); - _executor = new LLamaInteractExecutor(model); + _executor = new InteractiveExecutor(model); _executor.LoadState(executorStateFile); Console.WriteLine("Loaded state!"); } diff --git a/LLama.Examples/Program.cs b/LLama.Examples/Program.cs index f6827c0b..36e7ecc4 100644 --- a/LLama.Examples/Program.cs +++ b/LLama.Examples/Program.cs @@ -26,7 +26,8 @@ if(version == 1) Console.WriteLine("The examples for new versions are under working now. We'll soon update the examples." + " Thank you for your support!"); string modelPath = "D:\\development\\llama\\weights\\wizard-vicuna-13B.ggmlv3.q4_1.bin"; - var prompt = File.ReadAllText("Assets/chat-with-bob.txt").Trim(); + //var prompt = File.ReadAllText("Assets/chat-with-bob.txt").Trim(); + string prompt = " Qeustion: how to do binary search for an array in C#? Answer: "; //LLamaInteractExecutor ex = new(new LLamaModel(new ModelParams(modelPath, contextSize: 1024, seed: 1337))); @@ -39,22 +40,32 @@ if(version == 1) // prompt = Console.ReadLine(); //} - LLama.Examples.NewVersion.SaveAndLoadState runner = new(modelPath, prompt); + StatelessExecutor ex = new(new LLamaModel(new ModelParams(modelPath, contextSize: 256))); while (true) { - var input = Console.ReadLine(); - if(input == "save") + foreach (var text in ex.Infer(prompt, new SessionParams() { Temperature = 0.6f, AntiPrompts = new List { "user:" }, MaxTokens = 256 })) { - Console.Write("Your path to save state: "); - input = Console.ReadLine(); - runner.SaveState("./ex_state.json", input); - runner.LoadState("./ex_state.json", input); - } - else - { - runner.Run(input); + Console.Write(text); } + prompt = Console.ReadLine(); } + + //LLama.Examples.NewVersion.SaveAndLoadState runner = new(modelPath, prompt); + //while (true) + //{ + // var input = Console.ReadLine(); + // if(input == "save") + // { + // Console.Write("Your path to save state: "); + // input = Console.ReadLine(); + // runner.SaveState("./ex_state.json", input); + // runner.LoadState("./ex_state.json", input); + // } + // else + // { + // runner.Run(input); + // } + //} } else { diff --git a/LLama/Abstractions/Params/SessionParams.cs b/LLama/Abstractions/Params/SessionParams.cs index 444148f9..c846cefc 100644 --- a/LLama/Abstractions/Params/SessionParams.cs +++ b/LLama/Abstractions/Params/SessionParams.cs @@ -10,12 +10,12 @@ namespace LLama.Abstractions.Params /// /// number of tokens to keep from initial prompt /// - public int TokensToKeep { get; set; } = 0; + public int TokensKeep { get; set; } = 0; /// /// how many new tokens to predict (n_predict), set to -1 to inifinitely generate response /// until it complete. /// - public int ResponseTokensCount { get; set; } = -1; + public int MaxTokens { get; set; } = -1; /// /// logit bias for specific tokens /// diff --git a/LLama/ILLamaExecutor.cs b/LLama/ILLamaExecutor.cs index 89f9e45e..9e3a59ab 100644 --- a/LLama/ILLamaExecutor.cs +++ b/LLama/ILLamaExecutor.cs @@ -8,6 +8,7 @@ namespace LLama { public interface ILLamaExecutor { + public LLamaModel Model { get; } IEnumerable Infer(string text, SessionParams? sessionParams = null); IAsyncEnumerable InferAsync(string text, SessionParams? sessionParams = null, CancellationToken token = default); diff --git a/LLama/LLamaExecutorBase.cs b/LLama/LLamaExecutorBase.cs index b6f5f517..2d549169 100644 --- a/LLama/LLamaExecutorBase.cs +++ b/LLama/LLamaExecutorBase.cs @@ -15,9 +15,9 @@ using System.Threading.Tasks; namespace LLama { using llama_token = Int32; - public abstract class LLamaExecutorBase: ILLamaExecutor + public abstract class ChatExecutorBase: ILLamaExecutor { - protected LLamaModel _model; + protected readonly LLamaModel _model; protected int _pastTokensCount; // n_past protected int _consumedTokensCount; // n_consume protected int _n_session_consumed; @@ -28,7 +28,7 @@ namespace LLama protected List _session_tokens = new(); protected FixedSizeQuene _last_n_tokens; public LLamaModel Model => _model; - protected LLamaExecutorBase(LLamaModel model) + protected ChatExecutorBase(LLamaModel model) { _model = model; _pastTokensCount = 0; @@ -39,7 +39,7 @@ namespace LLama _last_n_tokens = new FixedSizeQuene(_model.ContextSize).FillWith(0); } - public unsafe LLamaExecutorBase WithSessionFile(string filename) + public unsafe ChatExecutorBase WithSessionFile(string filename) { _pathSession = filename; if (string.IsNullOrEmpty(filename)) @@ -129,7 +129,7 @@ namespace LLama InferStateArgs args = new InferStateArgs() { Antiprompts = sessionParams.AntiPrompts.ToList(), - RemainedTokens = sessionParams.ResponseTokensCount, + RemainedTokens = sessionParams.MaxTokens, ReturnValue = false, WaitForInput = false, NeedToSaveSession = !string.IsNullOrEmpty(_pathSession) && _n_matching_session_tokens < _embed_inps.Count @@ -177,7 +177,7 @@ namespace LLama InferStateArgs args = new InferStateArgs() { Antiprompts = sessionParams.AntiPrompts.ToList(), - RemainedTokens = sessionParams.ResponseTokensCount, + RemainedTokens = sessionParams.MaxTokens, ReturnValue = false, WaitForInput = false, NeedToSaveSession = !string.IsNullOrEmpty(_pathSession) && _n_matching_session_tokens < _embed_inps.Count diff --git a/LLama/LLamaInstructExecutor.cs b/LLama/LLamaInstructExecutor.cs index b3c25dd3..fa9e6ce4 100644 --- a/LLama/LLamaInstructExecutor.cs +++ b/LLama/LLamaInstructExecutor.cs @@ -12,12 +12,12 @@ using System.Text.Json.Serialization; namespace LLama { using llama_token = Int32; - public class LLamaInstructExecutor : LLamaExecutorBase + public class InstructExecutor : ChatExecutorBase { bool _is_prompt_run = true; llama_token[] _inp_pfx; llama_token[] _inp_sfx; - public LLamaInstructExecutor(LLamaModel model, string inputPrefix = "\n\n### Instruction:\n\n", + public InstructExecutor(LLamaModel model, string inputPrefix = "\n\n### Instruction:\n\n", string inputSuffix = "\n\n### Response:\n\n") : base(model) { _inp_pfx = _model.Tokenize(inputPrefix, true).ToArray(); @@ -131,9 +131,9 @@ namespace LLama args.WaitForInput = true; } - if (args.RemainedTokens <= 0 && sessionParams.ResponseTokensCount != -1) + if (args.RemainedTokens <= 0 && sessionParams.MaxTokens != -1) { - args.RemainedTokens = sessionParams.ResponseTokensCount; + args.RemainedTokens = sessionParams.MaxTokens; args.WaitForInput = true; } return false; @@ -145,7 +145,7 @@ namespace LLama _is_prompt_run = false; if (_pastTokensCount + _embeds.Count > _model.ContextSize) { - HandleRunOutOfContext(sessionParams.TokensToKeep); + HandleRunOutOfContext(sessionParams.TokensKeep); } TryReuseMathingPrefix(); @@ -162,19 +162,7 @@ namespace LLama if (_embed_inps.Count <= _consumedTokensCount && !args.WaitForInput) { - var temp = sessionParams.Temperature; - var top_k = sessionParams.TopK <= 0 ? NativeApi.llama_n_vocab(_model.NativeHandle) : sessionParams.TopK; - var top_p = sessionParams.TopK; - var tfs_z = sessionParams.TfsZ; - var typical_p = sessionParams.TypicalP; var repeat_last_n = sessionParams.RepeatLastTokensCount < 0 ? _model.ContextSize : sessionParams.RepeatLastTokensCount; - var repeat_penalty = sessionParams.RepeatPenalty; - var alpha_presence = sessionParams.PresencePenalty; - var alpha_frequency = sessionParams.FrequencyPenalty; - var mirostat = sessionParams.Mirostat; - var mirostat_tau = sessionParams.MirostatTau; - var mirostat_eta = sessionParams.MirostatEta; - var penalize_nl = sessionParams.PenalizeNL; // optionally save the session on first sample (for faster prompt loading next time) if (!string.IsNullOrEmpty(_pathSession) && args.NeedToSaveSession) @@ -184,10 +172,10 @@ namespace LLama } var tokenDataArray = _model.ApplyPenalty(_last_n_tokens, sessionParams.LogitBias, repeat_last_n, - repeat_penalty, alpha_frequency, alpha_presence, penalize_nl); + sessionParams.RepeatPenalty, sessionParams.FrequencyPenalty, sessionParams.PresencePenalty, sessionParams.PenalizeNL); - var id = _model.Sample(tokenDataArray, temp, mirostat, mirostat_tau, mirostat_eta, top_k, top_p, - tfs_z, typical_p); + var id = _model.Sample(tokenDataArray, sessionParams.Temperature, sessionParams.Mirostat, sessionParams.MirostatTau, + sessionParams.MirostatEta, sessionParams.TopK, sessionParams.TopP, sessionParams.TfsZ, sessionParams.TypicalP); _last_n_tokens.Enqueue(id); diff --git a/LLama/LLamaInteractExecutor.cs b/LLama/LLamaInteractExecutor.cs index edbf6800..45a4f8d7 100644 --- a/LLama/LLamaInteractExecutor.cs +++ b/LLama/LLamaInteractExecutor.cs @@ -15,11 +15,11 @@ using System.Threading.Tasks; namespace LLama { using llama_token = Int32; - public class LLamaInteractExecutor : LLamaExecutorBase + public class InteractiveExecutor : ChatExecutorBase { bool _is_prompt_run = true; llama_token[] _llama_token_newline; - public LLamaInteractExecutor(LLamaModel model) : base(model) + public InteractiveExecutor(LLamaModel model) : base(model) { _llama_token_newline = Utils.Tokenize(_model.NativeHandle, "\n", false, _model.Encoding).ToArray(); } @@ -134,9 +134,9 @@ namespace LLama return true; } - if (args.RemainedTokens <= 0 && sessionParams.ResponseTokensCount != -1) + if (args.RemainedTokens <= 0 && sessionParams.MaxTokens != -1) { - args.RemainedTokens = sessionParams.ResponseTokensCount; + args.RemainedTokens = sessionParams.MaxTokens; args.WaitForInput = true; } return false; @@ -149,7 +149,7 @@ namespace LLama _is_prompt_run = false; if (_pastTokensCount + _embeds.Count > _model.ContextSize) { - HandleRunOutOfContext(sessionParams.TokensToKeep); + HandleRunOutOfContext(sessionParams.TokensKeep); } TryReuseMathingPrefix(); @@ -166,19 +166,7 @@ namespace LLama if (_embed_inps.Count <= _consumedTokensCount && !args.WaitForInput) { - var temp = sessionParams.Temperature; - var top_k = sessionParams.TopK <= 0 ? NativeApi.llama_n_vocab(_model.NativeHandle) : sessionParams.TopK; - var top_p = sessionParams.TopK; - var tfs_z = sessionParams.TfsZ; - var typical_p = sessionParams.TypicalP; var repeat_last_n = sessionParams.RepeatLastTokensCount < 0 ? _model.ContextSize : sessionParams.RepeatLastTokensCount; - var repeat_penalty = sessionParams.RepeatPenalty; - var alpha_presence = sessionParams.PresencePenalty; - var alpha_frequency = sessionParams.FrequencyPenalty; - var mirostat = sessionParams.Mirostat; - var mirostat_tau = sessionParams.MirostatTau; - var mirostat_eta = sessionParams.MirostatEta; - var penalize_nl = sessionParams.PenalizeNL; // optionally save the session on first sample (for faster prompt loading next time) if (!string.IsNullOrEmpty(_pathSession) && args.NeedToSaveSession) @@ -188,10 +176,10 @@ namespace LLama } var tokenDataArray = _model.ApplyPenalty(_last_n_tokens, sessionParams.LogitBias, repeat_last_n, - repeat_penalty, alpha_frequency, alpha_presence, penalize_nl); + sessionParams.RepeatPenalty, sessionParams.FrequencyPenalty, sessionParams.PresencePenalty, sessionParams.PenalizeNL); - var id = _model.Sample(tokenDataArray, temp, mirostat, mirostat_tau, mirostat_eta, top_k, top_p, - tfs_z, typical_p); + var id = _model.Sample(tokenDataArray, sessionParams.Temperature, sessionParams.Mirostat, sessionParams.MirostatTau, + sessionParams.MirostatEta, sessionParams.TopK, sessionParams.TopP, sessionParams.TfsZ, sessionParams.TypicalP); _last_n_tokens.Enqueue(id); diff --git a/LLama/LLamaModel.cs b/LLama/LLamaModel.cs index db14b72a..5b3546bf 100644 --- a/LLama/LLamaModel.cs +++ b/LLama/LLamaModel.cs @@ -67,28 +67,46 @@ namespace LLama /// /// public void SaveState(string filename) + { + File.WriteAllBytes(filename, GetStateData()); + } + + /// + /// Get the state data as a byte array. + /// + /// + public byte[] GetStateData() { var stateSize = NativeApi.llama_get_state_size(_ctx); byte[] stateMemory = new byte[stateSize]; NativeApi.llama_copy_state_data(_ctx, stateMemory); - File.WriteAllBytes(filename, stateMemory); + return stateMemory; } /// /// Load the state from specified path. /// /// - /// Whether to clear previous footprints of this model. /// - public void LoadState(string filename, bool clearPreviousEmbed = true) + public void LoadState(string filename) { var stateMemory = File.ReadAllBytes(filename); + LoadState(stateMemory); + } + + /// + /// Load the state from memory. + /// + /// + /// + public void LoadState(byte[] stateData) + { int stateSize = (int)NativeApi.llama_get_state_size(_ctx); - if (stateMemory.Length != stateSize) + if (stateData.Length != stateSize) { throw new RuntimeError("Failed to validate state size."); } - NativeApi.llama_set_state_data(_ctx, stateMemory); + NativeApi.llama_set_state_data(_ctx, stateData); } public llama_token Sample(LLamaTokenDataArray candidates, float temperature = 0.8f, MiroStateType mirostat = MiroStateType.Disable, diff --git a/LLama/LLamaStatelessExecutor.cs b/LLama/LLamaStatelessExecutor.cs new file mode 100644 index 00000000..1ba79e68 --- /dev/null +++ b/LLama/LLamaStatelessExecutor.cs @@ -0,0 +1,116 @@ +using LLama.Abstractions.Params; +using LLama.Native; +using System; +using System.Collections.Generic; +using System.Diagnostics.CodeAnalysis; +using System.Linq; +using System.Runtime.CompilerServices; +using System.Text; +using System.Threading; + +namespace LLama +{ + using llama_token = Int32; + /// + /// This executor infer the input as one-time job. Previous inputs won't impact on the + /// response to current input. + /// + public class StatelessExecutor : ILLamaExecutor + { + private LLamaModel _model; + private byte[] _originalState; + public LLamaModel Model => _model; + public StatelessExecutor(LLamaModel model) + { + _model = model; + var tokens = model.Tokenize(" ", true); + Utils.Eval(_model.NativeHandle, tokens.ToArray(), 0, tokens.Count(), 0, _model.Params.Threads); + _originalState = model.GetStateData(); + } + public IEnumerable Infer(string text, SessionParams? sessionParams = null) + { + int n_past = 1; + if(sessionParams is null) + { + sessionParams = new SessionParams(); + } + List lastTokens = new(sessionParams.RepeatLastTokensCount); + for(int i = 0; i < lastTokens.Count; i++) + { + lastTokens[i] = 0; + } + List tokens = _model.Tokenize(text, true).ToList(); + int n_prompt_tokens = tokens.Count; + + Utils.Eval(_model.NativeHandle, tokens.ToArray(), 0, n_prompt_tokens, n_past, _model.Params.Threads); + + lastTokens.AddRange(tokens); + n_past += n_prompt_tokens; + + int max_tokens = sessionParams.MaxTokens < 0 ? int.MaxValue : sessionParams.MaxTokens; + for(int i = 0; i < max_tokens; i++) + { + var repeat_last_n = sessionParams.RepeatLastTokensCount < 0 ? _model.ContextSize : sessionParams.RepeatLastTokensCount; + + var tokenDataArray = _model.ApplyPenalty(lastTokens, sessionParams.LogitBias, repeat_last_n, + sessionParams.RepeatPenalty, sessionParams.FrequencyPenalty, sessionParams.PresencePenalty, sessionParams.PenalizeNL); + + var id = _model.Sample(tokenDataArray, sessionParams.Temperature, sessionParams.Mirostat, sessionParams.MirostatTau, + sessionParams.MirostatEta, sessionParams.TopK, sessionParams.TopP, sessionParams.TfsZ, sessionParams.TypicalP); + + lastTokens.Add(id); + + string response = Utils.TokenToString(id, _model.NativeHandle, _model.Encoding); + yield return response; + + tokens.Clear(); + tokens.Add(id); + + if (sessionParams.AntiPrompts is not null && sessionParams.AntiPrompts.Count() > 0) + { + string last_output = ""; + foreach (var token in lastTokens) + { + last_output += Utils.PtrToString(NativeApi.llama_token_to_str(_model.NativeHandle, id), _model.Encoding); + } + + bool should_break = false; + foreach (var antiprompt in sessionParams.AntiPrompts) + { + if (last_output.EndsWith(antiprompt)) + { + should_break = true; + break; + } + } + if (should_break) + { + break; + } + } + + // when run out of context + if (n_past + tokens.Count > _model.ContextSize) + { + int n_left = n_past - sessionParams.TokensKeep; + + n_past = Math.Max(1, sessionParams.TokensKeep); + + // insert n_left/2 tokens at the start of embed from last_n_tokens + tokens.InsertRange(0, lastTokens.Take(lastTokens.Count - tokens.Count).Skip(_model.ContextSize - n_left / 2 - tokens.Count)); + } + + n_past = _model.Eval(tokens.ToArray(), n_past); + } + + _model.LoadState(_originalState); + } + + + public async IAsyncEnumerable InferAsync(string text, SessionParams? sessionParams = null, [EnumeratorCancellation] CancellationToken token = default) + { + yield return ""; + throw new NotImplementedException(); + } + } +}