BotSharp/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Embedding/TextEmbeddingProvider.cs
2024-10-02 13:00:23 -05:00

75 lines
2.1 KiB
C#

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;
private const int DEFAULT_DIMENSION = 1536;
protected string _model;
protected int _dimension;
public virtual string Provider => "azure-openai";
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.ToFloats().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.ToFloats().ToArray()).ToList();
}
public void SetModelName(string model)
{
_model = model;
}
public void SetDimension(int dimension)
{
_dimension = dimension > 0 ? dimension : DEFAULT_DIMENSION;
}
public int GetDimension()
{
return _dimension;
}
private EmbeddingGenerationOptions PrepareOptions()
{
return new EmbeddingGenerationOptions
{
Dimensions = GetDimensionOption()
};
}
private int GetDimensionOption()
{
return _dimension > 0 ? _dimension : DEFAULT_DIMENSION;
}
}