BotSharp/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Embedding/TextEmbeddingProvider.cs

74 lines
2.3 KiB
C#
Raw Normal View History

2024-07-01 02:14:02 +00:00
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;
2024-08-22 15:15:05 +00:00
private const int DEFAULT_DIMENSION = 3072;
2024-07-01 02:14:02 +00:00
protected string _model;
public virtual string Provider => "azure-openai";
2024-07-01 22:05:32 +00:00
public int Dimension { get; set; }
2024-07-01 02:14:02 +00:00
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 state = _services.GetRequiredService<IConversationStateService>();
2024-07-01 22:05:32 +00:00
var stateDimension = state.GetState("embedding_dimension");
2024-07-03 02:51:44 +00:00
var defaultDimension = Dimension > 0 ? Dimension : DEFAULT_DIMENSION;
2024-07-01 02:14:02 +00:00
if (int.TryParse(stateDimension, out var dimension))
{
2024-07-03 02:51:44 +00:00
return dimension > 0 ? dimension : defaultDimension;
2024-07-01 02:14:02 +00:00
}
2024-07-03 02:51:44 +00:00
return defaultDimension;
2024-07-01 02:14:02 +00:00
}
}