From 6cc9f14b8c2fd0d98f11f1ab67780874b7d0f317 Mon Sep 17 00:00:00 2001 From: Jicheng Lu <103353@smsassist.com> Date: Mon, 24 Jun 2024 13:03:05 -0500 Subject: [PATCH] add image generation --- .../MLTasks/IChatCompletion.cs | 3 + .../MLTasks/Settings/LlmModelSetting.cs | 3 +- .../Controllers/InstructModeController.cs | 32 ++++++++++ .../Instructs/ImageGenerationViewModel.cs | 18 ++++++ .../Providers/ChatCompletionProvider.cs | 5 ++ .../Providers/ChatCompletionProvider.cs | 58 +++++++++++++++++++ .../Providers/ChatCompletionProvider.cs | 5 ++ .../Providers/ChatCompletionProvider.cs | 5 ++ .../Providers/ChatCompletionProvider.cs | 5 ++ .../Providers/ChatCompletionProvider.cs | 5 ++ .../SemanticKernelChatCompletionProvider.cs | 5 ++ .../Providers/ChatCompletionProvider.cs | 7 ++- 12 files changed, 148 insertions(+), 3 deletions(-) create mode 100644 src/Infrastructure/BotSharp.OpenAPI/ViewModels/Instructs/ImageGenerationViewModel.cs diff --git a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IChatCompletion.cs b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IChatCompletion.cs index 306e72e3..c5332b40 100644 --- a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IChatCompletion.cs +++ b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IChatCompletion.cs @@ -24,4 +24,7 @@ public interface IChatCompletion Task GetChatCompletionsStreamingAsync(Agent agent, List conversations, Func onMessageReceived); + + Task GetImageGeneration(Agent agent, + List conversations); } diff --git a/src/Infrastructure/BotSharp.Abstraction/MLTasks/Settings/LlmModelSetting.cs b/src/Infrastructure/BotSharp.Abstraction/MLTasks/Settings/LlmModelSetting.cs index b86578fe..0a0bd42a 100644 --- a/src/Infrastructure/BotSharp.Abstraction/MLTasks/Settings/LlmModelSetting.cs +++ b/src/Infrastructure/BotSharp.Abstraction/MLTasks/Settings/LlmModelSetting.cs @@ -51,5 +51,6 @@ public class LlmModelSetting public enum LlmModelType { Text = 1, - Chat = 2 + Chat = 2, + Image = 3 } diff --git a/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs b/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs index 45120a26..0b013456 100644 --- a/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs +++ b/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs @@ -103,4 +103,36 @@ public class InstructModeController : ControllerBase return $"Error in analyzing files."; } } + + [HttpPost("/instruct/image-generation")] + public async Task ImageGeneration([FromBody] IncomingMessageModel input) + { + var state = _services.GetRequiredService(); + input.States.ForEach(x => state.SetState(x.Key, x.Value, activeRounds: x.ActiveRounds, source: StateSource.External)); + var imageViewModel = new ImageGenerationViewModel(); + + try + { + var completion = CompletionProvider.GetChatCompletion(_services, provider: input.Provider ?? "openai", + modelId: input.ModelId ?? "dall-e"); + var message = await completion.GetImageGeneration(new Agent() + { + Id = Guid.Empty.ToString(), + }, new List + { + new RoleDialogModel(AgentRole.User, input.Text) + }); + + imageViewModel.Content = message.Content; + imageViewModel.Data = message.Data; + return imageViewModel; + } + catch (Exception ex) + { + var error = "Error in image generation."; + _logger.LogError($"{error} {ex.Message}"); + imageViewModel.Message = error; + return imageViewModel; + } + } } diff --git a/src/Infrastructure/BotSharp.OpenAPI/ViewModels/Instructs/ImageGenerationViewModel.cs b/src/Infrastructure/BotSharp.OpenAPI/ViewModels/Instructs/ImageGenerationViewModel.cs new file mode 100644 index 00000000..179c05f4 --- /dev/null +++ b/src/Infrastructure/BotSharp.OpenAPI/ViewModels/Instructs/ImageGenerationViewModel.cs @@ -0,0 +1,18 @@ +using System.Text.Json.Serialization; + +namespace BotSharp.OpenAPI.ViewModels.Instructs; + +public class ImageGenerationViewModel +{ + [JsonPropertyName("content")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Content { get; set; } + + [JsonPropertyName("data")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public object? Data { get; set; } + + [JsonPropertyName("message")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Message { get; set; } +} diff --git a/src/Plugins/BotSharp.Plugin.AnthropicAI/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.AnthropicAI/Providers/ChatCompletionProvider.cs index 87202109..8977e04d 100644 --- a/src/Plugins/BotSharp.Plugin.AnthropicAI/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.AnthropicAI/Providers/ChatCompletionProvider.cs @@ -264,4 +264,9 @@ public class ChatCompletionProvider : IChatCompletion { _model = model; } + + public Task GetImageGeneration(Agent agent, List conversations) + { + throw new NotImplementedException(); + } } diff --git a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs index 50fd33ce..d85b4d65 100644 --- a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/ChatCompletionProvider.cs @@ -446,4 +446,62 @@ public class ChatCompletionProvider : IChatCompletion functionResultData = "31 celsius"; return new ChatRequestToolMessage(functionResultData.ToString(), toolCall.Id); } + + public async Task GetImageGeneration(Agent agent, List conversations) + { + var contentHooks = _services.GetServices().ToList(); + foreach (var hook in contentHooks) + { + await hook.BeforeGenerating(agent, conversations); + } + + var client = ProviderHelper.GetClient(Provider, _model, _services); + var options = BuildImageGenerationOptions(conversations); + var response = await client.GetImageGenerationsAsync(options); + var image = response.Value.Data.First(); + + var content = string.Empty; + if (!string.IsNullOrEmpty(image.RevisedPrompt)) + { + content = image.RevisedPrompt; + } + + var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) + { + CurrentAgentId = agent.Id, + MessageId = conversations.LastOrDefault()?.MessageId ?? string.Empty, + Data = image.Url.AbsoluteUri ?? image.Base64Data + }; + + foreach (var hook in contentHooks) + { + await hook.AfterGenerated(responseMessage, new TokenStatsModel + { + Prompt = options.Prompt, + Provider = Provider, + Model = _model, + PromptCount = options.Prompt.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count(), + CompletionCount = content.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count() + }); + } + + return responseMessage; + } + + private ImageGenerationOptions BuildImageGenerationOptions(List conversations) + { + var state = _services.GetRequiredService(); + + var sizeValue = !string.IsNullOrEmpty(state.GetState("image_size")) ? state.GetState("image_size") : "1024x1024"; + var qualityValue = !string.IsNullOrEmpty(state.GetState("image_quality")) ? state.GetState("image_quality") : "standard"; + + var options = new ImageGenerationOptions + { + DeploymentName = _model, + Prompt = conversations.LastOrDefault()?.Payload ?? conversations.LastOrDefault()?.Content ?? string.Empty, + Size = new ImageSize(sizeValue), + Quality = new ImageGenerationQuality(qualityValue) + }; + return options; + } } diff --git a/src/Plugins/BotSharp.Plugin.GoogleAI/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.GoogleAI/Providers/ChatCompletionProvider.cs index d278b110..c4628c6d 100644 --- a/src/Plugins/BotSharp.Plugin.GoogleAI/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.GoogleAI/Providers/ChatCompletionProvider.cs @@ -149,4 +149,9 @@ public class ChatCompletionProvider : IChatCompletion { _model = model; } + + public async Task GetImageGeneration(Agent agent, List conversations) + { + throw new NotImplementedException(); + } } diff --git a/src/Plugins/BotSharp.Plugin.HuggingFace/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.HuggingFace/Providers/ChatCompletionProvider.cs index 88b37dbb..80c950c8 100644 --- a/src/Plugins/BotSharp.Plugin.HuggingFace/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.HuggingFace/Providers/ChatCompletionProvider.cs @@ -139,4 +139,9 @@ public class ChatCompletionProvider : IChatCompletion return msg; } + + public async Task GetImageGeneration(Agent agent, List conversations) + { + throw new NotImplementedException(); + } } diff --git a/src/Plugins/BotSharp.Plugin.LLamaSharp/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.LLamaSharp/Providers/ChatCompletionProvider.cs index 0db444ce..cdd96a0b 100644 --- a/src/Plugins/BotSharp.Plugin.LLamaSharp/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.LLamaSharp/Providers/ChatCompletionProvider.cs @@ -191,4 +191,9 @@ public class ChatCompletionProvider : IChatCompletion { _model = model; } + + public async Task GetImageGeneration(Agent agent, List conversations) + { + throw new NotImplementedException(); + } } diff --git a/src/Plugins/BotSharp.Plugin.MetaGLM/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.MetaGLM/Providers/ChatCompletionProvider.cs index c702a7ec..7010842d 100644 --- a/src/Plugins/BotSharp.Plugin.MetaGLM/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.MetaGLM/Providers/ChatCompletionProvider.cs @@ -231,6 +231,11 @@ public class ChatCompletionProvider : IChatCompletion throw new NotImplementedException(); } + public async Task GetImageGeneration(Agent agent, List conversations) + { + throw new NotImplementedException(); + } + public void SetModelName(string model) { _model = model; diff --git a/src/Plugins/BotSharp.Plugin.SemanticKernel/SemanticKernelChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.SemanticKernel/SemanticKernelChatCompletionProvider.cs index 156f238c..cb70b189 100644 --- a/src/Plugins/BotSharp.Plugin.SemanticKernel/SemanticKernelChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.SemanticKernel/SemanticKernelChatCompletionProvider.cs @@ -102,5 +102,10 @@ namespace BotSharp.Plugin.SemanticKernel { _model = model; } + + public async Task GetImageGeneration(Agent agent, List conversations) + { + throw new NotImplementedException(); + } } } \ No newline at end of file diff --git a/src/Plugins/BotSharp.Plugin.SparkDesk/Providers/ChatCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.SparkDesk/Providers/ChatCompletionProvider.cs index 7d529df1..25e588c3 100644 --- a/src/Plugins/BotSharp.Plugin.SparkDesk/Providers/ChatCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.SparkDesk/Providers/ChatCompletionProvider.cs @@ -1,8 +1,6 @@ using BotSharp.Abstraction.Agents; using BotSharp.Abstraction.Agents.Enums; using BotSharp.Abstraction.Loggers; -using BotSharp.Abstraction.Routing; -using Sdcb.SparkDesk.ResponseInternals; namespace BotSharp.Plugin.SparkDesk.Providers; @@ -270,4 +268,9 @@ public class ChatCompletionProvider : IChatCompletion FunctionDef functionDef = new FunctionDef(def.Name, def.Description, fundef.ToArray()); return functionDef; } + + public async Task GetImageGeneration(Agent agent, List conversations) + { + throw new NotImplementedException(); + } }