add text embedding
This commit is contained in:
parent
9d25ea627a
commit
5fe7a5f9d3
|
|
@ -6,7 +6,8 @@ public interface ITextEmbedding
|
|||
/// The Embedding provider like Microsoft Azure, OpenAI, ClaudAI
|
||||
/// </summary>
|
||||
string Provider { get; }
|
||||
int Dimension { get; }
|
||||
int Dimension { get; set; }
|
||||
Task<float[]> GetVectorAsync(string text);
|
||||
Task<List<float[]>> GetVectorsAsync(List<string> texts);
|
||||
void SetModelName(string model);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -57,5 +57,6 @@ public enum LlmModelType
|
|||
{
|
||||
Text = 1,
|
||||
Chat = 2,
|
||||
Image = 3
|
||||
Image = 3,
|
||||
Embedding = 4
|
||||
}
|
||||
|
|
|
|||
|
|
@ -99,6 +99,26 @@ public class CompletionProvider
|
|||
return completer;
|
||||
}
|
||||
|
||||
public static ITextEmbedding GetTextEmbedding(IServiceProvider services,
|
||||
string? provider = null,
|
||||
string? model = null)
|
||||
{
|
||||
var completions = services.GetServices<ITextEmbedding>();
|
||||
|
||||
(provider, model) = GetProviderAndModel(services, provider: provider, model: model);
|
||||
|
||||
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;
|
||||
}
|
||||
|
||||
private static (string, string) GetProviderAndModel(IServiceProvider services,
|
||||
string? provider = null,
|
||||
string? model = null,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
<Project Sdk="Microsoft.NET.Sdk">
|
||||
<Project Sdk="Microsoft.NET.Sdk">
|
||||
|
||||
<PropertyGroup>
|
||||
<TargetFramework>$(TargetFramework)</TargetFramework>
|
||||
|
|
|
|||
|
|
@ -0,0 +1,41 @@
|
|||
using BotSharp.Core.Infrastructures;
|
||||
using BotSharp.OpenAPI.ViewModels.Embeddings;
|
||||
|
||||
namespace BotSharp.OpenAPI.Controllers;
|
||||
|
||||
[Authorize]
|
||||
[ApiController]
|
||||
public class TextEmbeddingController : ControllerBase
|
||||
{
|
||||
private readonly IServiceProvider _services;
|
||||
private readonly ILogger<TextEmbeddingController> _logger;
|
||||
|
||||
public TextEmbeddingController(
|
||||
IServiceProvider services,
|
||||
ILogger<TextEmbeddingController> logger)
|
||||
{
|
||||
_services = services;
|
||||
_logger = logger;
|
||||
}
|
||||
|
||||
[HttpPost("/text-embedding/generation")]
|
||||
public async Task<List<float[]>> GenerateTextEmbeddings(EmbeddingInputModel input)
|
||||
{
|
||||
var state = _services.GetRequiredService<IConversationStateService>();
|
||||
input.States.ForEach(x => state.SetState(x.Key, x.Value, activeRounds: x.ActiveRounds, source: StateSource.External));
|
||||
|
||||
try
|
||||
{
|
||||
var completion = CompletionProvider.GetTextEmbedding(_services, provider: input.Provider ?? "openai", model: input.Model ?? "text-embedding-3-small");
|
||||
completion.Dimension = input.Dimension;
|
||||
|
||||
var embeddings = await completion.GetVectorsAsync(input.Texts?.ToList() ?? []);
|
||||
return embeddings;
|
||||
}
|
||||
catch (Exception ex)
|
||||
{
|
||||
_logger.LogError($"Error when generating text embeddings... {ex.Message}");
|
||||
throw;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,12 @@
|
|||
using System.Text.Json.Serialization;
|
||||
|
||||
namespace BotSharp.OpenAPI.ViewModels.Embeddings;
|
||||
|
||||
public class EmbeddingInputModel : MessageConfig
|
||||
{
|
||||
[JsonPropertyName("texts")]
|
||||
public IEnumerable<string> Texts { get; set; } = new List<string>();
|
||||
|
||||
[JsonPropertyName("dimension")]
|
||||
public int Dimension { get; set; }
|
||||
}
|
||||
|
|
@ -1,6 +1,7 @@
|
|||
using BotSharp.Abstraction.Plugins;
|
||||
using BotSharp.Abstraction.Settings;
|
||||
using BotSharp.Plugin.AzureOpenAI.Providers.Chat;
|
||||
using BotSharp.Plugin.AzureOpenAI.Providers.Embedding;
|
||||
using BotSharp.Plugin.AzureOpenAI.Providers.Image;
|
||||
using BotSharp.Plugin.AzureOpenAI.Providers.Text;
|
||||
using Microsoft.Extensions.Configuration;
|
||||
|
|
@ -31,5 +32,7 @@ public class AzureOpenAiPlugin : IBotSharpPlugin
|
|||
services.AddScoped<IChatCompletion, OpenAiChatCompletionProvider>();
|
||||
services.AddScoped<IImageGeneration, ImageGenerationProvider>();
|
||||
services.AddScoped<IImageGeneration, OpenAiImageGenerationProvider>();
|
||||
services.AddScoped<ITextEmbedding, TextEmbeddingProvider>();
|
||||
services.AddScoped<ITextEmbedding, OpenAiTextEmbeddingProvider>();
|
||||
}
|
||||
}
|
||||
|
|
@ -6,7 +6,7 @@ public class ChatCompletionProvider : IChatCompletion
|
|||
{
|
||||
protected readonly AzureOpenAiSettings _settings;
|
||||
protected readonly IServiceProvider _services;
|
||||
protected readonly ILogger _logger;
|
||||
protected readonly ILogger<ChatCompletionProvider> _logger;
|
||||
|
||||
protected string _model;
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,10 @@
|
|||
namespace BotSharp.Plugin.AzureOpenAI.Providers.Embedding;
|
||||
|
||||
public class OpenAiTextEmbeddingProvider : TextEmbeddingProvider
|
||||
{
|
||||
public override string Provider => "openai";
|
||||
|
||||
public OpenAiTextEmbeddingProvider(AzureOpenAiSettings settings,
|
||||
ILogger<OpenAiTextEmbeddingProvider> logger,
|
||||
IServiceProvider services) : base(settings, logger, services) { }
|
||||
}
|
||||
|
|
@ -0,0 +1,71 @@
|
|||
using OpenAI.Embeddings;
|
||||
|
||||
namespace BotSharp.Plugin.AzureOpenAI.Providers.Embedding;
|
||||
|
||||
public class TextEmbeddingProvider : ITextEmbedding
|
||||
{
|
||||
protected readonly AzureOpenAiSettings _settings;
|
||||
protected readonly IServiceProvider _services;
|
||||
protected readonly ILogger<TextEmbeddingProvider> _logger;
|
||||
|
||||
protected string _model;
|
||||
|
||||
public virtual string Provider => "azure-openai";
|
||||
|
||||
public int Dimension { get; set; } = 4096;
|
||||
|
||||
public TextEmbeddingProvider(
|
||||
AzureOpenAiSettings settings,
|
||||
ILogger<TextEmbeddingProvider> logger,
|
||||
IServiceProvider services)
|
||||
{
|
||||
_settings = settings;
|
||||
_logger = logger;
|
||||
_services = services;
|
||||
}
|
||||
|
||||
public async Task<float[]> GetVectorAsync(string text)
|
||||
{
|
||||
var client = ProviderHelper.GetClient(Provider, _model, _services);
|
||||
var embeddingClient = client.GetEmbeddingClient(_model);
|
||||
var options = PrepareOptions();
|
||||
var response = await embeddingClient.GenerateEmbeddingAsync(text, options);
|
||||
var value = response.Value;
|
||||
return value.Vector.ToArray();
|
||||
}
|
||||
|
||||
public async Task<List<float[]>> GetVectorsAsync(List<string> texts)
|
||||
{
|
||||
var client = ProviderHelper.GetClient(Provider, _model, _services);
|
||||
var embeddingClient = client.GetEmbeddingClient(_model);
|
||||
var options = PrepareOptions();
|
||||
var response = await embeddingClient.GenerateEmbeddingsAsync(texts, options);
|
||||
var value = response.Value;
|
||||
return value.Select(x => x.Vector.ToArray()).ToList();
|
||||
}
|
||||
|
||||
public void SetModelName(string model)
|
||||
{
|
||||
_model = model;
|
||||
}
|
||||
|
||||
private EmbeddingGenerationOptions PrepareOptions()
|
||||
{
|
||||
return new EmbeddingGenerationOptions
|
||||
{
|
||||
Dimensions = GetDimension()
|
||||
};
|
||||
}
|
||||
|
||||
private int GetDimension()
|
||||
{
|
||||
var defaultDimension = 4096;
|
||||
var state = _services.GetRequiredService<IConversationStateService>();
|
||||
var stateDimension = state.GetState("embedding_state");
|
||||
if (int.TryParse(stateDimension, out var dimension))
|
||||
{
|
||||
return dimension > 0 ? dimension : defaultDimension;
|
||||
}
|
||||
return Dimension > 0 ? Dimension : defaultDimension;
|
||||
}
|
||||
}
|
||||
|
|
@ -7,7 +7,7 @@ public class TextEmbeddingProvider : ITextEmbedding
|
|||
private LLamaEmbedder _embedder;
|
||||
private readonly LlamaSharpSettings _settings;
|
||||
private readonly IServiceProvider _services;
|
||||
public int Dimension => 4096;
|
||||
public int Dimension { get; set; } = 4096;
|
||||
|
||||
public string Provider => "llama-sharp";
|
||||
|
||||
|
|
@ -34,4 +34,6 @@ public class TextEmbeddingProvider : ITextEmbedding
|
|||
{
|
||||
throw new NotImplementedException();
|
||||
}
|
||||
|
||||
public void SetModelName(string model) { }
|
||||
}
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ public class fastTextEmbeddingProvider : ITextEmbedding
|
|||
private FastTextWrapper _fastText;
|
||||
private readonly IServiceProvider _services;
|
||||
|
||||
private int dimension;
|
||||
public int Dimension
|
||||
{
|
||||
get
|
||||
|
|
@ -21,6 +22,10 @@ public class fastTextEmbeddingProvider : ITextEmbedding
|
|||
LoadModel();
|
||||
return _fastText.GetModelDimension();
|
||||
}
|
||||
set
|
||||
{
|
||||
dimension = value;
|
||||
}
|
||||
}
|
||||
|
||||
public string Provider => "meta-ai";
|
||||
|
|
@ -66,4 +71,6 @@ public class fastTextEmbeddingProvider : ITextEmbedding
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
public void SetModelName(string model) { }
|
||||
}
|
||||
|
|
|
|||
|
|
@ -49,5 +49,7 @@ namespace BotSharp.Plugin.SemanticKernel
|
|||
return embeddings.Select(_ => _.ToArray())
|
||||
.ToList();
|
||||
}
|
||||
|
||||
public void SetModelName(string model) { }
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue