From 7244fd2604b40f6ec42a729bfe668ef4f677f773 Mon Sep 17 00:00:00 2001 From: Haiping Chen Date: Tue, 27 Jun 2023 14:17:53 -0500 Subject: [PATCH] Enable LLamaSharp (WIP) #81 --- .../Services/ConversationService.cs | 2 +- .../Knowledges/Services/KnowledgeService.cs | 2 +- .../LLamaSharp/ChatCompletionProvider.cs | 78 ++++--------------- .../Plugins/LLamaSharp/LLamaSharpPlugin.cs | 2 + .../Plugins/LLamaSharp/LlamaAiModel.cs | 32 ++++++++ .../LLamaSharp/TextEmbeddingProvider.cs | 19 +++++ .../Providers/ChatCompletionProvider.cs | 2 +- .../Providers/TextCompletionProvider.cs | 2 +- src/WebStarter/appsettings.json | 3 +- 9 files changed, 75 insertions(+), 67 deletions(-) create mode 100644 src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LlamaAiModel.cs create mode 100644 src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/TextEmbeddingProvider.cs diff --git a/src/Infrastructure/BotSharp.Core/Conversations/Services/ConversationService.cs b/src/Infrastructure/BotSharp.Core/Conversations/Services/ConversationService.cs index aaac8c91..4363c678 100644 --- a/src/Infrastructure/BotSharp.Core/Conversations/Services/ConversationService.cs +++ b/src/Infrastructure/BotSharp.Core/Conversations/Services/ConversationService.cs @@ -85,7 +85,7 @@ public class ConversationService : IConversationService public IChatCompletion GetChatCompletion() { var completions = _services.GetServices(); - return completions.FirstOrDefault(x => x.GetType().FullName.Contains(_settings.ChatCompletion)); + return completions.FirstOrDefault(x => x.GetType().FullName.EndsWith(_settings.ChatCompletion)); } public Task CleanHistory(string agentId) diff --git a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs b/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs index 2c2d28ff..b0459dc4 100644 --- a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs +++ b/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs @@ -93,7 +93,7 @@ public class KnowledgeService : IKnowledgeService public ITextCompletion GetTextCompletion() { var textCompletion = _services.GetServices() - .FirstOrDefault(x => x.GetType().Name == _settings.TextCompletion); + .FirstOrDefault(x => x.GetType().FullName.EndsWith(_settings.TextCompletion)); return textCompletion; } } diff --git a/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/ChatCompletionProvider.cs b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/ChatCompletionProvider.cs index 915f9608..55b623c5 100644 --- a/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/ChatCompletionProvider.cs +++ b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/ChatCompletionProvider.cs @@ -8,34 +8,24 @@ namespace BotSharp.Core.Plugins.LLamaSharp; public class ChatCompletionProvider : IChatCompletion { - private readonly IChatModel _model; - private readonly LlamaSharpSettings _settings; + private IChatModel _model; - public ChatCompletionProvider(LlamaSharpSettings settings) + + public ChatCompletionProvider(LlamaAiModel model) { - _settings = settings; - _model = new LLamaModel(new LLamaParams(model: _settings.ModelPath, - n_ctx: _settings.MaxContextLength, - interactive: _settings.Interactive, - repeat_penalty: _settings.RepeatPenalty, - verbose_prompt: _settings.VerbosePrompt, - n_gpu_layers: _settings.NumberOfGpuLayer)); - - var prompt = GetInstruction(); - _model.InitChatPrompt(prompt, "UTF-8"); - _model.InitChatAntiprompt(new string[] { "user:" }); + model.LoadModel(); + _model = model.Model; + // _model.InitChatPrompt(prompt, "UTF-8"); + // _model.InitChatAntiprompt(new string[] { "user:" }); } - public int Priority => 100; - public async Task GetChatCompletionsAsync(List conversations, Func onChunkReceived) { string totalResponse = ""; - var prompt = GetInstruction(); var content = string.Join("\n ", conversations.Select(x => $"{x.Role}: {x.Text.Replace("user:", "")}")).Trim(); content += "\n assistant: "; - foreach (var response in _model.Chat(content, prompt, "UTF-8")) + foreach (var response in _model.Chat(content, "", "UTF-8")) { Console.Write(response); totalResponse += response; @@ -48,54 +38,18 @@ public class ChatCompletionProvider : IChatCompletion public Task GetChatCompletionsAsync(Agent agent, List conversations) { string totalResponse = ""; - var prompt = GetInstruction(); - var content = string.Join("\n ", conversations.Select(x => $"{x.Role}: {x.Text.Replace("user:", "")}")).Trim(); - content += "\n assistant: "; - foreach (var response in _model.Chat(content, prompt, "UTF-8")) + var content = string.Join("\n", conversations.Select(x => $"{x.Role}: {x.Text.Replace("user:", "")}")).Trim(); + content += "\nassistant: "; + foreach (var response in _model.Chat(content, agent.Instruction, "UTF-8")) { + if (response == "\n") + { + break; + } Console.Write(response); totalResponse += response; } - return Task.FromResult(totalResponse); - } - - public List GetChatSamples() - { - var samples = new List(); - if (!string.IsNullOrEmpty(_settings.ChatSampleFile)) - { - var lines = File.ReadAllLines(_settings.ChatSampleFile); - for (int i = 0; i < lines.Length; i++) - { - var line = lines[i]; - var role = line.Substring(0, line.IndexOf(' ') - 1); - var content = line.Substring(line.IndexOf(' ') + 1); - - samples.Add(new RoleDialogModel - { - Role = role, - Text = content - }); - } - } - return samples; - } - - public string GetInstruction() - { - var instruction = ""; - if (!string.IsNullOrEmpty(_settings.InstructionFile)) - { - instruction = File.ReadAllText(_settings.InstructionFile); - } - - instruction += "\n"; - foreach (var message in GetChatSamples()) - { - instruction += $"\n{message.Role}: {message.Text}"; - } - - return instruction; + return Task.FromResult(totalResponse.Trim()); } } diff --git a/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LLamaSharpPlugin.cs b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LLamaSharpPlugin.cs index 0a882cff..7a2470d8 100644 --- a/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LLamaSharpPlugin.cs +++ b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LLamaSharpPlugin.cs @@ -11,6 +11,8 @@ public class LLamaSharpPlugin : IBotSharpPlugin config.Bind("LlamaSharp", llamaSharpSettings); services.AddSingleton(x => llamaSharpSettings); + services.AddSingleton(); + services.AddScoped(); services.AddScoped(); } } diff --git a/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LlamaAiModel.cs b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LlamaAiModel.cs new file mode 100644 index 00000000..3ee01589 --- /dev/null +++ b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/LlamaAiModel.cs @@ -0,0 +1,32 @@ +using LLama; +namespace BotSharp.Core.Plugins.LLamaSharp; + +public class LlamaAiModel +{ + private readonly LlamaSharpSettings _settings; + + LLamaModel _model; + + public LLamaModel Model => _model; + + public LlamaAiModel(LlamaSharpSettings settings) + { + _settings = settings; + } + + + public void LoadModel() + { + if (_model != null) + { + return; + } + + _model = new LLamaModel(new LLamaParams(model: _settings.ModelPath, + n_ctx: _settings.MaxContextLength, + interactive: _settings.Interactive, + repeat_penalty: _settings.RepeatPenalty, + verbose_prompt: _settings.VerbosePrompt, + n_gpu_layers: _settings.NumberOfGpuLayer)); + } +} diff --git a/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/TextEmbeddingProvider.cs b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/TextEmbeddingProvider.cs new file mode 100644 index 00000000..fa6e956a --- /dev/null +++ b/src/Infrastructure/BotSharp.Core/Plugins/LLamaSharp/TextEmbeddingProvider.cs @@ -0,0 +1,19 @@ +using BotSharp.Abstraction.MLTasks; + +namespace BotSharp.Core.Plugins.LLamaSharp; + +public class TextEmbeddingProvider : ITextEmbedding +{ + public int Dimension => throw new NotImplementedException(); + private readonly LlamaAiModel _llama; + + public TextEmbeddingProvider(LlamaAiModel llama) + { + _llama = llama; + } + + public float[] GetVector(string text) + { + return new float[0]; + } +} diff --git a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs index ee0a1359..b7fa4c7a 100644 --- a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs @@ -88,7 +88,7 @@ public class ChatCompletionProvider : IChatCompletion } } - return output; + return output.Trim(); } private ChatCompletionsOptions PrepareOptions(Agent agent, List conversations) diff --git a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/TextCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/TextCompletionProvider.cs index 8ee9bfbc..85fb2d25 100644 --- a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/TextCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/TextCompletionProvider.cs @@ -41,7 +41,7 @@ public class TextCompletionProvider : ITextCompletion completion += t.Text; }; - return completion; + return completion.Trim(); } private OpenAIClient GetOpenAIClient() diff --git a/src/WebStarter/appsettings.json b/src/WebStarter/appsettings.json index 4c479f32..bc17abf3 100644 --- a/src/WebStarter/appsettings.json +++ b/src/WebStarter/appsettings.json @@ -15,6 +15,7 @@ "Conversation": { "ChatCompletion": "AzureOpenAI.Providers.ChatCompletionProvider" + // "ChatCompletion": "LLamaSharp.ChatCompletionProvider" }, "LlamaSharp": { @@ -85,7 +86,7 @@ "Plugins": [ "KnowledgeBasePlugin", "MemVecDbPlugin", - // "LLamaSharpPlugin", + "LLamaSharpPlugin", "AzureOpenAiPlugin", "MetaAiPlugin", "QdrantPlugin",