Merge pull request #554 from iceljc/features/add-image-edit

Features/add image edit
This commit is contained in:
iceljc 2024-07-18 23:11:30 -05:00 committed by GitHub
commit 37fbd6d636
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
18 changed files with 457 additions and 675 deletions

View file

@ -2,7 +2,7 @@ using System.IO;
namespace BotSharp.Abstraction.MLTasks;
public interface IImageVariation
public interface IImageCompletion
{
/// <summary>
/// The LLM provider like Microsoft Azure, OpenAI, ClaudAI
@ -15,5 +15,7 @@ public interface IImageVariation
/// <param name="model">deployment name</param>
void SetModelName(string model);
Task<RoleDialogModel> GetImageGeneration(Agent agent, RoleDialogModel message);
Task<RoleDialogModel> GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName);
}

View file

@ -1,5 +0,0 @@
namespace BotSharp.Abstraction.MLTasks;
public interface IImageEdit
{
}

View file

@ -1,17 +0,0 @@
namespace BotSharp.Abstraction.MLTasks;
public interface IImageGeneration
{
/// <summary>
/// The LLM provider like Microsoft Azure, OpenAI, ClaudAI
/// </summary>
string Provider { get; }
/// <summary>
/// Set model name, one provider can consume different model or version(s)
/// </summary>
/// <param name="model">deployment name</param>
void SetModelName(string model);
Task<RoleDialogModel> GetImageGeneration(Agent agent, RoleDialogModel message);
}

View file

@ -6,7 +6,7 @@ public partial class BotSharpFileService
{
public async Task<RoleDialogModel> GenerateImage(string? provider, string? model, string text)
{
var completion = CompletionProvider.GetImageGeneration(_services, provider: provider ?? "openai", model: model ?? "dall-e-3");
var completion = CompletionProvider.GetImageCompletion(_services, provider: provider ?? "openai", model: model ?? "dall-e-3");
var message = await completion.GetImageGeneration(new Agent()
{
Id = Guid.Empty.ToString(),
@ -21,7 +21,7 @@ public partial class BotSharpFileService
throw new ArgumentException($"Please fill in at least file url or file data!");
}
var completion = CompletionProvider.GetImageVariation(_services, provider: provider ?? "openai", model: model ?? "dall-e-2");
var completion = CompletionProvider.GetImageCompletion(_services, provider: provider ?? "openai", model: model ?? "dall-e-2");
var bytes = await DownloadFile(file);
using var stream = new MemoryStream();
stream.Write(bytes, 0, bytes.Length);

View file

@ -5,7 +5,7 @@ namespace BotSharp.Core.Infrastructures;
public class CompletionProvider
{
public static object? GetCompletion(IServiceProvider services,
public static object GetCompletion(IServiceProvider services,
string? provider = null,
string? model = null,
AgentLlmConfig? agentConfig = null)
@ -18,26 +18,20 @@ public class CompletionProvider
if (settings.Type == LlmModelType.Text)
{
return GetTextCompletion(services,
provider: provider,
model: model,
agentConfig: agentConfig);
return GetTextCompletion(services, provider: provider, model: model, agentConfig: agentConfig);
}
else if (settings.Type == LlmModelType.Embedding)
{
return GetTextEmbedding(services,
provider: provider,
model: model);
return GetTextEmbedding(services, provider: provider, model: model);
}
else if (settings.Type == LlmModelType.Chat)
else if (settings.Type == LlmModelType.Image)
{
return GetChatCompletion(services,
provider: provider,
model: model,
agentConfig: agentConfig);
return GetImageCompletion(services, provider: provider, model: model);
}
else
{
return GetChatCompletion(services, provider: provider, model: model, agentConfig: agentConfig);
}
return null;
}
public static IChatCompletion GetChatCompletion(IServiceProvider services,
@ -82,36 +76,15 @@ public class CompletionProvider
return completer;
}
public static IImageGeneration GetImageGeneration(IServiceProvider services,
string? provider = null,
string? model = null,
string? modelId = null,
bool imageGenerate = false,
AgentLlmConfig? agentConfig = null)
{
var completions = services.GetServices<IImageGeneration>();
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, modelId: modelId,
imageGenerate: imageGenerate, agentConfig: agentConfig);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)
{
var logger = services.GetRequiredService<ILogger<CompletionProvider>>();
logger.LogError($"Can't resolve completion provider by {provider}");
}
completer?.SetModelName(model);
return completer;
}
public static IImageVariation GetImageVariation(IServiceProvider services,
public static IImageCompletion GetImageCompletion(IServiceProvider services,
string? provider = null,
string? model = null,
string? modelId = null,
bool imageGenerate = false)
{
var completions = services.GetServices<IImageVariation>();
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, modelId: modelId, imageGenerate: imageGenerate);
var completions = services.GetServices<IImageCompletion>();
(provider, model) = GetProviderAndModel(services, provider: provider,
model: model, modelId: modelId, imageGenerate: imageGenerate);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)

View file

@ -141,6 +141,10 @@ public class InstructModeController : ControllerBase
try
{
var file = input.Files.FirstOrDefault(x => !string.IsNullOrWhiteSpace(x.FileUrl) || !string.IsNullOrWhiteSpace(x.FileData));
if (file == null)
{
return new ImageGenerationViewModel { Message = "Error! Cannot find an image!" };
}
var message = await fileService.VarifyImage(input.Provider, input.Model, file);
imageViewModel.Content = message.Content;
imageViewModel.Images = message.GeneratedImages.Select(x => ImageViewModel.ToViewModel(x)).ToList();

View file

@ -29,7 +29,6 @@ public class AzureOpenAiPlugin : IBotSharpPlugin
services.AddScoped<ITextCompletion, TextCompletionProvider>();
services.AddScoped<IChatCompletion, ChatCompletionProvider>();
services.AddScoped<ITextEmbedding, TextEmbeddingProvider>();
services.AddScoped<IImageGeneration, ImageGenerationProvider>();
services.AddScoped<IImageVariation, ImageVariationProvider>();
services.AddScoped<IImageCompletion, ImageCompletionProvider>();
}
}

View file

@ -0,0 +1,71 @@
using OpenAI.Images;
namespace BotSharp.Plugin.AzureOpenAI.Providers.Image;
public partial class ImageCompletionProvider
{
public async Task<RoleDialogModel> GetImageGeneration(Agent agent, RoleDialogModel message)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (prompt, imageCount, options) = PrepareOptions(message);
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImages(prompt, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
return await Task.FromResult(responseMessage);
}
private (string, int, ImageGenerationOptions) PrepareOptions(RoleDialogModel message)
{
var prompt = message?.Payload ?? message?.Content ?? string.Empty;
var state = _services.GetRequiredService<IConversationStateService>();
var size = GetImageSize(state.GetState("image_size"));
var quality = GetImageQuality(state.GetState("image_quality"));
var style = GetImageStyle(state.GetState("image_style"));
var format = GetImageFormat(state.GetState("image_format"));
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageGenerationOptions
{
Size = size,
Quality = quality,
Style = style,
ResponseFormat = format
};
return (prompt, count, options);
}
}

View file

@ -0,0 +1,65 @@
using OpenAI.Images;
namespace BotSharp.Plugin.AzureOpenAI.Providers.Image;
public partial class ImageCompletionProvider
{
public async Task<RoleDialogModel> GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (imageCount, options) = PrepareOptions();
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImageVariations(image, imageFileName, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
return await Task.FromResult(responseMessage);
}
private (int, ImageVariationOptions) PrepareOptions()
{
var state = _services.GetRequiredService<IConversationStateService>();
var size = GetImageSize(state.GetState("image_size"));
var format = GetImageFormat(state.GetState("image_format"));
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageVariationOptions
{
Size = size,
ResponseFormat = format
};
return (count, options);
}
}

View file

@ -2,11 +2,11 @@ using OpenAI.Images;
namespace BotSharp.Plugin.AzureOpenAI.Providers.Image;
public class ImageGenerationProvider : IImageGeneration
public partial class ImageCompletionProvider : IImageCompletion
{
protected readonly AzureOpenAiSettings _settings;
protected readonly IServiceProvider _services;
protected readonly ILogger _logger;
protected readonly ILogger<ImageCompletionProvider> _logger;
private const int DEFAULT_IMAGE_COUNT = 1;
private const int IMAGE_COUNT_LIMIT = 5;
@ -15,9 +15,9 @@ public class ImageGenerationProvider : IImageGeneration
public virtual string Provider => "azure-openai";
public ImageGenerationProvider(
public ImageCompletionProvider(
AzureOpenAiSettings settings,
ILogger<ImageGenerationProvider> logger,
ILogger<ImageCompletionProvider> logger,
IServiceProvider services)
{
_settings = settings;
@ -25,91 +25,12 @@ public class ImageGenerationProvider : IImageGeneration
_logger = logger;
}
public async Task<RoleDialogModel> GetImageGeneration(Agent agent, RoleDialogModel message)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (prompt, imageCount, options) = PrepareOptions(message);
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImages(prompt, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
// After
var contentHooks = _services.GetServices<IContentGeneratingHook>().ToList();
foreach (var hook in contentHooks)
{
await hook.AfterGenerated(responseMessage, new TokenStatsModel
{
Prompt = prompt,
Provider = Provider,
Model = _model,
PromptCount = prompt.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count(),
CompletionCount = content.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count()
});
}
return responseMessage;
}
public void SetModelName(string model)
{
_model = model;
}
private (string, int, ImageGenerationOptions) PrepareOptions(RoleDialogModel message)
{
var prompt = message?.Payload ?? message?.Content ?? string.Empty;
var state = _services.GetRequiredService<IConversationStateService>();
var size = state.GetState("image_size");
var quality = state.GetState("image_quality");
var style = state.GetState("image_style");
var format = state.GetState("image_format");
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageGenerationOptions
{
Size = GetImageSize(size),
Quality = GetImageQuality(quality),
Style = GetImageStyle(style),
ResponseFormat = GetImageFormat(format)
};
return (prompt, count, options);
}
#region Private methods
private GeneratedImageSize GetImageSize(string size)
{
var value = !string.IsNullOrEmpty(size) ? size : "1024x1024";
@ -210,6 +131,15 @@ public class ImageGenerationProvider : IImageGeneration
return DEFAULT_IMAGE_COUNT;
}
return retCount > 0 && retCount <= IMAGE_COUNT_LIMIT ? retCount : DEFAULT_IMAGE_COUNT;
if (retCount <= 0)
{
retCount = DEFAULT_IMAGE_COUNT;
}
else if (retCount > IMAGE_COUNT_LIMIT)
{
retCount = IMAGE_COUNT_LIMIT;
}
return retCount;
}
#endregion
}

View file

@ -1,152 +0,0 @@
using OpenAI.Images;
namespace BotSharp.Plugin.AzureOpenAI.Providers.Image;
public class ImageVariationProvider : IImageVariation
{
protected readonly AzureOpenAiSettings _settings;
protected readonly IServiceProvider _services;
protected readonly ILogger<ImageVariationProvider> _logger;
private const int DEFAULT_IMAGE_COUNT = 1;
private const int IMAGE_COUNT_LIMIT = 5;
protected string _model;
public virtual string Provider => "azure-openai";
public ImageVariationProvider(
AzureOpenAiSettings settings,
ILogger<ImageVariationProvider> logger,
IServiceProvider services)
{
_settings = settings;
_services = services;
_logger = logger;
}
public async Task<RoleDialogModel> GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (imageCount, options) = PrepareOptions();
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImageVariations(image, imageFileName, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
return await Task.FromResult(responseMessage);
}
public void SetModelName(string model)
{
_model = model;
}
private (int, ImageVariationOptions) PrepareOptions()
{
var state = _services.GetRequiredService<IConversationStateService>();
var size = state.GetState("image_size");
var format = state.GetState("image_format");
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageVariationOptions
{
Size = GetImageSize(size),
ResponseFormat = GetImageFormat(format)
};
return (count, options);
}
private GeneratedImageSize GetImageSize(string size)
{
var value = !string.IsNullOrEmpty(size) ? size : "1024x1024";
GeneratedImageSize retSize;
switch (value)
{
case "256x256":
retSize = GeneratedImageSize.W256xH256;
break;
case "512x512":
retSize = GeneratedImageSize.W512xH512;
break;
case "1024x1024":
retSize = GeneratedImageSize.W1024xH1024;
break;
case "1024x1792":
retSize = GeneratedImageSize.W1024xH1792;
break;
case "1792x1024":
retSize = GeneratedImageSize.W1792xH1024;
break;
default:
retSize = GeneratedImageSize.W1024xH1024;
break;
}
return retSize;
}
private GeneratedImageFormat GetImageFormat(string format)
{
var value = !string.IsNullOrEmpty(format) ? format : "uri";
GeneratedImageFormat retFormat;
switch (value)
{
case "uri":
retFormat = GeneratedImageFormat.Uri;
break;
case "bytes":
retFormat = GeneratedImageFormat.Bytes;
break;
default:
retFormat = GeneratedImageFormat.Uri;
break;
}
return retFormat;
}
private int GetImageCount(string count)
{
if (!int.TryParse(count, out var retCount))
{
return DEFAULT_IMAGE_COUNT;
}
return retCount > 0 && retCount <= IMAGE_COUNT_LIMIT ? retCount : DEFAULT_IMAGE_COUNT;
}
}

View file

@ -18,7 +18,6 @@ public class GenerateImageFn : IFunctionCallback
_logger = logger;
}
public async Task<bool> Execute(RoleDialogModel message)
{
var args = JsonSerializer.Deserialize<LlmFileContext>(message.FunctionArgs);
@ -59,7 +58,7 @@ public class GenerateImageFn : IFunctionCallback
{
try
{
var completion = CompletionProvider.GetImageGeneration(_services, provider: "openai", model: "dall-e-3");
var completion = CompletionProvider.GetImageCompletion(_services, provider: "openai", model: "dall-e-3");
var text = !string.IsNullOrWhiteSpace(description) ? description : message.Content;
var dialog = RoleDialogModel.From(message, AgentRole.User, text);
var result = await completion.GetImageGeneration(agent, dialog);

View file

@ -29,7 +29,6 @@ public class OpenAiPlugin : IBotSharpPlugin
services.AddScoped<ITextCompletion, TextCompletionProvider>();
services.AddScoped<IChatCompletion, ChatCompletionProvider>();
services.AddScoped<ITextEmbedding, TextEmbeddingProvider>();
services.AddScoped<IImageGeneration, ImageGenerationProvider>();
services.AddScoped<IImageVariation, ImageVariationProvider>();
services.AddScoped<IImageCompletion, ImageCompletionProvider>();
}
}

View file

@ -0,0 +1,71 @@
using OpenAI.Images;
namespace BotSharp.Plugin.OpenAI.Providers.Image;
public partial class ImageCompletionProvider
{
public async Task<RoleDialogModel> GetImageGeneration(Agent agent, RoleDialogModel message)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (prompt, imageCount, options) = PrepareOptions(message);
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImages(prompt, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
return await Task.FromResult(responseMessage);
}
private (string, int, ImageGenerationOptions) PrepareOptions(RoleDialogModel message)
{
var prompt = message?.Payload ?? message?.Content ?? string.Empty;
var state = _services.GetRequiredService<IConversationStateService>();
var size = GetImageSize(state.GetState("image_size"));
var quality = GetImageQuality(state.GetState("image_quality"));
var style = GetImageStyle(state.GetState("image_style"));
var format = GetImageFormat(state.GetState("image_format"));
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageGenerationOptions
{
Size = size,
Quality = quality,
Style = style,
ResponseFormat = format
};
return (prompt, count, options);
}
}

View file

@ -0,0 +1,65 @@
using OpenAI.Images;
namespace BotSharp.Plugin.OpenAI.Providers.Image;
public partial class ImageCompletionProvider
{
public async Task<RoleDialogModel> GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (imageCount, options) = PrepareOptions();
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImageVariations(image, imageFileName, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
return await Task.FromResult(responseMessage);
}
private (int, ImageVariationOptions) PrepareOptions()
{
var state = _services.GetRequiredService<IConversationStateService>();
var size = GetImageSize(state.GetState("image_size"));
var format = GetImageFormat(state.GetState("image_format"));
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageVariationOptions
{
Size = size,
ResponseFormat = format
};
return (count, options);
}
}

View file

@ -0,0 +1,145 @@
using OpenAI.Images;
namespace BotSharp.Plugin.OpenAI.Providers.Image;
public partial class ImageCompletionProvider : IImageCompletion
{
protected readonly OpenAiSettings _settings;
protected readonly IServiceProvider _services;
protected readonly ILogger<ImageCompletionProvider> _logger;
private const int DEFAULT_IMAGE_COUNT = 1;
private const int IMAGE_COUNT_LIMIT = 5;
protected string _model;
public virtual string Provider => "openai";
public ImageCompletionProvider(
OpenAiSettings settings,
ILogger<ImageCompletionProvider> logger,
IServiceProvider services)
{
_settings = settings;
_services = services;
_logger = logger;
}
public void SetModelName(string model)
{
_model = model;
}
#region Private methods
private GeneratedImageSize GetImageSize(string size)
{
var value = !string.IsNullOrEmpty(size) ? size : "1024x1024";
GeneratedImageSize retSize;
switch (value)
{
case "256x256":
retSize = GeneratedImageSize.W256xH256;
break;
case "512x512":
retSize = GeneratedImageSize.W512xH512;
break;
case "1024x1024":
retSize = GeneratedImageSize.W1024xH1024;
break;
case "1024x1792":
retSize = GeneratedImageSize.W1024xH1792;
break;
case "1792x1024":
retSize = GeneratedImageSize.W1792xH1024;
break;
default:
retSize = GeneratedImageSize.W1024xH1024;
break;
}
return retSize;
}
private GeneratedImageQuality GetImageQuality(string quality)
{
var value = !string.IsNullOrEmpty(quality) ? quality : "standard";
GeneratedImageQuality retQuality;
switch (value)
{
case "standard":
retQuality = GeneratedImageQuality.Standard;
break;
case "hd":
retQuality = GeneratedImageQuality.High;
break;
default:
retQuality = GeneratedImageQuality.Standard;
break;
}
return retQuality;
}
private GeneratedImageStyle GetImageStyle(string style)
{
var value = !string.IsNullOrEmpty(style) ? style : "natural";
GeneratedImageStyle retStyle;
switch (value)
{
case "natural":
retStyle = GeneratedImageStyle.Natural;
break;
case "vivid":
retStyle = GeneratedImageStyle.Vivid;
break;
default:
retStyle = GeneratedImageStyle.Natural;
break;
}
return retStyle;
}
private GeneratedImageFormat GetImageFormat(string format)
{
var value = !string.IsNullOrEmpty(format) ? format : "uri";
GeneratedImageFormat retFormat;
switch (value)
{
case "uri":
retFormat = GeneratedImageFormat.Uri;
break;
case "bytes":
retFormat = GeneratedImageFormat.Bytes;
break;
default:
retFormat = GeneratedImageFormat.Uri;
break;
}
return retFormat;
}
private int GetImageCount(string count)
{
if (!int.TryParse(count, out var retCount))
{
return DEFAULT_IMAGE_COUNT;
}
if (retCount <= 0)
{
retCount = DEFAULT_IMAGE_COUNT;
}
else if (retCount > IMAGE_COUNT_LIMIT)
{
retCount = IMAGE_COUNT_LIMIT;
}
return retCount;
}
#endregion
}

View file

@ -1,215 +0,0 @@
using OpenAI.Images;
namespace BotSharp.Plugin.OpenAI.Providers.Image;
public class ImageGenerationProvider : IImageGeneration
{
protected readonly OpenAiSettings _settings;
protected readonly IServiceProvider _services;
protected readonly ILogger<ImageGenerationProvider> _logger;
private const int DEFAULT_IMAGE_COUNT = 1;
private const int IMAGE_COUNT_LIMIT = 5;
protected string _model;
public virtual string Provider => "openai";
public ImageGenerationProvider(
OpenAiSettings settings,
ILogger<ImageGenerationProvider> logger,
IServiceProvider services)
{
_settings = settings;
_services = services;
_logger = logger;
}
public async Task<RoleDialogModel> GetImageGeneration(Agent agent, RoleDialogModel message)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (prompt, imageCount, options) = PrepareOptions(message);
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImages(prompt, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
// After
var contentHooks = _services.GetServices<IContentGeneratingHook>().ToList();
foreach (var hook in contentHooks)
{
await hook.AfterGenerated(responseMessage, new TokenStatsModel
{
Prompt = prompt,
Provider = Provider,
Model = _model,
PromptCount = prompt.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count(),
CompletionCount = content.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count()
});
}
return responseMessage;
}
public void SetModelName(string model)
{
_model = model;
}
private (string, int, ImageGenerationOptions) PrepareOptions(RoleDialogModel message)
{
var prompt = message?.Payload ?? message?.Content ?? string.Empty;
var state = _services.GetRequiredService<IConversationStateService>();
var size = state.GetState("image_size");
var quality = state.GetState("image_quality");
var style = state.GetState("image_style");
var format = state.GetState("image_format");
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageGenerationOptions
{
Size = GetImageSize(size),
Quality = GetImageQuality(quality),
Style = GetImageStyle(style),
ResponseFormat = GetImageFormat(format)
};
return (prompt, count, options);
}
private GeneratedImageSize GetImageSize(string size)
{
var value = !string.IsNullOrEmpty(size) ? size : "1024x1024";
GeneratedImageSize retSize;
switch (value)
{
case "256x256":
retSize = GeneratedImageSize.W256xH256;
break;
case "512x512":
retSize = GeneratedImageSize.W512xH512;
break;
case "1024x1024":
retSize = GeneratedImageSize.W1024xH1024;
break;
case "1024x1792":
retSize = GeneratedImageSize.W1024xH1792;
break;
case "1792x1024":
retSize = GeneratedImageSize.W1792xH1024;
break;
default:
retSize = GeneratedImageSize.W1024xH1024;
break;
}
return retSize;
}
private GeneratedImageQuality GetImageQuality(string quality)
{
var value = !string.IsNullOrEmpty(quality) ? quality : "standard";
GeneratedImageQuality retQuality;
switch (value)
{
case "standard":
retQuality = GeneratedImageQuality.Standard;
break;
case "hd":
retQuality = GeneratedImageQuality.High;
break;
default:
retQuality = GeneratedImageQuality.Standard;
break;
}
return retQuality;
}
private GeneratedImageStyle GetImageStyle(string style)
{
var value = !string.IsNullOrEmpty(style) ? style : "natural";
GeneratedImageStyle retStyle;
switch (value)
{
case "natural":
retStyle = GeneratedImageStyle.Natural;
break;
case "vivid":
retStyle = GeneratedImageStyle.Vivid;
break;
default:
retStyle = GeneratedImageStyle.Natural;
break;
}
return retStyle;
}
private GeneratedImageFormat GetImageFormat(string format)
{
var value = !string.IsNullOrEmpty(format) ? format : "uri";
GeneratedImageFormat retFormat;
switch (value)
{
case "uri":
retFormat = GeneratedImageFormat.Uri;
break;
case "bytes":
retFormat = GeneratedImageFormat.Bytes;
break;
default:
retFormat = GeneratedImageFormat.Uri;
break;
}
return retFormat;
}
private int GetImageCount(string count)
{
if (!int.TryParse(count, out var retCount))
{
return DEFAULT_IMAGE_COUNT;
}
return retCount > 0 && retCount <= IMAGE_COUNT_LIMIT ? retCount : DEFAULT_IMAGE_COUNT;
}
}

View file

@ -1,152 +0,0 @@
using OpenAI.Images;
namespace BotSharp.Plugin.OpenAI.Providers.Image;
public class ImageVariationProvider : IImageVariation
{
protected readonly OpenAiSettings _settings;
protected readonly IServiceProvider _services;
protected readonly ILogger<ImageVariationProvider> _logger;
private const int DEFAULT_IMAGE_COUNT = 1;
private const int IMAGE_COUNT_LIMIT = 5;
protected string _model;
public virtual string Provider => "openai";
public ImageVariationProvider(
OpenAiSettings settings,
ILogger<ImageVariationProvider> logger,
IServiceProvider services)
{
_settings = settings;
_services = services;
_logger = logger;
}
public async Task<RoleDialogModel> GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName)
{
var client = ProviderHelper.GetClient(Provider, _model, _services);
var (imageCount, options) = PrepareOptions();
var imageClient = client.GetImageClient(_model);
var response = imageClient.GenerateImageVariations(image, imageFileName, imageCount, options);
var values = response.Value;
var generatedImages = new List<ImageGeneration>();
foreach (var value in values)
{
if (value == null) continue;
var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty };
if (options.ResponseFormat == GeneratedImageFormat.Uri)
{
generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty;
}
else if (options.ResponseFormat == GeneratedImageFormat.Bytes)
{
var base64Str = string.Empty;
var bytes = value?.ImageBytes?.ToArray();
if (!bytes.IsNullOrEmpty())
{
base64Str = Convert.ToBase64String(bytes);
}
generatedImage.ImageData = base64Str;
}
generatedImages.Add(generatedImage);
}
var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description));
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = message?.MessageId ?? string.Empty,
GeneratedImages = generatedImages
};
return await Task.FromResult(responseMessage);
}
public void SetModelName(string model)
{
_model = model;
}
private (int, ImageVariationOptions) PrepareOptions()
{
var state = _services.GetRequiredService<IConversationStateService>();
var size = state.GetState("image_size");
var format = state.GetState("image_format");
var count = GetImageCount(state.GetState("image_count", "1"));
var options = new ImageVariationOptions
{
Size = GetImageSize(size),
ResponseFormat = GetImageFormat(format)
};
return (count, options);
}
private GeneratedImageSize GetImageSize(string size)
{
var value = !string.IsNullOrEmpty(size) ? size : "1024x1024";
GeneratedImageSize retSize;
switch (value)
{
case "256x256":
retSize = GeneratedImageSize.W256xH256;
break;
case "512x512":
retSize = GeneratedImageSize.W512xH512;
break;
case "1024x1024":
retSize = GeneratedImageSize.W1024xH1024;
break;
case "1024x1792":
retSize = GeneratedImageSize.W1024xH1792;
break;
case "1792x1024":
retSize = GeneratedImageSize.W1792xH1024;
break;
default:
retSize = GeneratedImageSize.W1024xH1024;
break;
}
return retSize;
}
private GeneratedImageFormat GetImageFormat(string format)
{
var value = !string.IsNullOrEmpty(format) ? format : "uri";
GeneratedImageFormat retFormat;
switch (value)
{
case "uri":
retFormat = GeneratedImageFormat.Uri;
break;
case "bytes":
retFormat = GeneratedImageFormat.Bytes;
break;
default:
retFormat = GeneratedImageFormat.Uri;
break;
}
return retFormat;
}
private int GetImageCount(string count)
{
if (!int.TryParse(count, out var retCount))
{
return DEFAULT_IMAGE_COUNT;
}
return retCount > 0 && retCount <= IMAGE_COUNT_LIMIT ? retCount : DEFAULT_IMAGE_COUNT;
}
}