add text embedding

This commit is contained in:
Jicheng Lu 2024-06-30 21:14:02 -05:00
parent 9d25ea627a
commit 5fe7a5f9d3
13 changed files with 175 additions and 5 deletions

View file

@ -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);
}

View file

@ -57,5 +57,6 @@ public enum LlmModelType
{
Text = 1,
Chat = 2,
Image = 3
Image = 3,
Embedding = 4
}

View file

@ -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,

View file

@ -1,4 +1,4 @@
<Project Sdk="Microsoft.NET.Sdk">
<Project Sdk="Microsoft.NET.Sdk">
<PropertyGroup>
<TargetFramework>$(TargetFramework)</TargetFramework>

View file

@ -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;
}
}
}

View file

@ -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; }
}

View file

@ -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>();
}
}

View file

@ -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;

View file

@ -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) { }
}

View file

@ -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;
}
}

View file

@ -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) { }
}

View file

@ -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) { }
}

View file

@ -49,5 +49,7 @@ namespace BotSharp.Plugin.SemanticKernel
return embeddings.Select(_ => _.ToArray())
.ToList();
}
public void SetModelName(string model) { }
}
}