refine knowledge settings

This commit is contained in:
Jicheng Lu 2024-08-22 15:23:22 -05:00
parent 4e83813c0a
commit fff9438294
20 changed files with 161 additions and 81 deletions

View file

@ -4,14 +4,28 @@ namespace BotSharp.Abstraction.Knowledges.Settings;
public class KnowledgeBaseSettings
{
public string DefaultCollection { get; set; } = KnowledgeCollectionName.BotSharp;
public string VectorDb { get; set; }
public string GraphDb { get; set; }
public KnowledgeModelSetting TextEmbedding { get; set; }
public DefaultKnowledgeBaseSetting Default { get; set; }
public List<VectorCollectionSetting> Collections { get; set; } = new();
}
public class KnowledgeModelSetting
public class DefaultKnowledgeBaseSetting
{
public string CollectionName { get; set; } = KnowledgeCollectionName.BotSharp;
public KnowledgeTextEmbeddingSetting TextEmbedding { get; set; }
}
public class VectorCollectionSetting
{
public string Name { get; set; }
public KnowledgeTextEmbeddingSetting TextEmbedding { get; set; }
}
public class KnowledgeTextEmbeddingSetting
{
public string Provider { get; set; }
public string Model { get; set; }
public int Dimension { get; set; }
}

View file

@ -6,8 +6,9 @@ public interface ITextEmbedding
/// The Embedding provider like Microsoft Azure, OpenAI, ClaudAI
/// </summary>
string Provider { get; }
int Dimension { get; set; }
Task<float[]> GetVectorAsync(string text);
Task<List<float[]>> GetVectorsAsync(List<string> texts);
void SetModelName(string model);
void SetDimension(int dimension);
int GetDimension();
}

View file

@ -52,6 +52,11 @@ public class LlmModelSetting
/// </summary>
public float CompletionCost { get; set; }
/// <summary>
/// Embedding dimension
/// </summary>
public int Dimension { get; set; }
public override string ToString()
{
return $"[{Type}] {Name} {Endpoint}";

View file

@ -111,7 +111,12 @@ public class CompletionProvider
logger.LogError($"Can't resolve completion provider by {provider}");
}
var llmProviderService = services.GetRequiredService<ILlmProviderService>();
var found = llmProviderService.GetSetting(provider, model);
completer.SetModelName(model);
completer.SetDimension(found.Dimension);
return completer;
}

View file

@ -18,7 +18,7 @@ public class TextEmbeddingController : ControllerBase
_logger = logger;
}
[HttpPost("/text-embedding/generation")]
[HttpPost("/text-embedding/generate")]
public async Task<List<float[]>> GenerateTextEmbeddings(EmbeddingInputModel input)
{
var state = _services.GetRequiredService<IConversationStateService>();
@ -27,7 +27,10 @@ public class TextEmbeddingController : ControllerBase
try
{
var completion = CompletionProvider.GetTextEmbedding(_services, provider: input.Provider ?? "openai", model: input.Model ?? "text-embedding-3-large");
completion.Dimension = input.Dimension;
if (input.Dimension.HasValue && input.Dimension.Value > 0)
{
completion.SetDimension(input.Dimension.Value);
}
var embeddings = await completion.GetVectorsAsync(input.Texts?.ToList() ?? []);
return embeddings;

View file

@ -8,6 +8,6 @@ public class EmbeddingInputModel : MessageConfig
public IEnumerable<string> Texts { get; set; } = new List<string>();
[JsonPropertyName("dimension")]
public int Dimension { get; set; } = 3072;
public int? Dimension { get; set; } = 3072;
}

View file

@ -10,11 +10,10 @@ public class TextEmbeddingProvider : ITextEmbedding
private const int DEFAULT_DIMENSION = 3072;
protected string _model;
protected int _dimension;
public virtual string Provider => "azure-openai";
public int Dimension { get; set; }
public TextEmbeddingProvider(
AzureOpenAiSettings settings,
ILogger<TextEmbeddingProvider> logger,
@ -50,24 +49,26 @@ public class TextEmbeddingProvider : ITextEmbedding
_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 = GetDimension()
Dimensions = GetDimensionOption()
};
}
private int GetDimension()
private int GetDimensionOption()
{
var state = _services.GetRequiredService<IConversationStateService>();
var stateDimension = state.GetState("embedding_dimension");
var defaultDimension = Dimension > 0 ? Dimension : DEFAULT_DIMENSION;
if (int.TryParse(stateDimension, out var dimension))
{
return dimension > 0 ? dimension : defaultDimension;
}
return defaultDimension;
return _dimension > 0 ? _dimension : DEFAULT_DIMENSION;
}
}

View file

@ -17,12 +17,11 @@ public class KnowledgeRetrievalFn : IFunctionCallback
{
var args = JsonSerializer.Deserialize<ExtractedKnowledge>(message.FunctionArgs ?? "{}");
var embedding = _services.GetServices<ITextEmbedding>().FirstOrDefault(x => x.Provider == _settings.TextEmbedding.Provider);
embedding.SetModelName(_settings.TextEmbedding.Model);
var collectionName = _settings.Default.CollectionName ?? KnowledgeCollectionName.BotSharp;
var embedding = KnowledgeSettingUtility.GetTextEmbeddingSetting(_services, collectionName);
var vector = await embedding.GetVectorAsync(args.Question);
var vectorDb = _services.GetServices<IVectorDb>().FirstOrDefault(x => x.Name == _settings.VectorDb);
var collectionName = !string.IsNullOrWhiteSpace(_settings.DefaultCollection) ? _settings.DefaultCollection : KnowledgeCollectionName.BotSharp;
var knowledges = await vectorDb.Search(collectionName, vector, new List<string> { KnowledgePayloadName.Text, KnowledgePayloadName.Answer });
if (!knowledges.IsNullOrEmpty())

View file

@ -17,8 +17,8 @@ public class MemorizeKnowledgeFn : IFunctionCallback
{
var args = JsonSerializer.Deserialize<ExtractedKnowledge>(message.FunctionArgs ?? "{}");
var embedding = _services.GetServices<ITextEmbedding>().FirstOrDefault(x => x.Provider == _settings.TextEmbedding.Provider);
embedding.SetModelName(_settings.TextEmbedding.Model);
var collectionName = _settings.Default.CollectionName ?? KnowledgeCollectionName.BotSharp;
var embedding = KnowledgeSettingUtility.GetTextEmbeddingSetting(_services, collectionName);
var vector = await embedding.GetVectorsAsync(new List<string>
{
@ -26,7 +26,6 @@ public class MemorizeKnowledgeFn : IFunctionCallback
});
var vectorDb = _services.GetServices<IVectorDb>().FirstOrDefault(x => x.Name == _settings.VectorDb);
var collectionName = !string.IsNullOrWhiteSpace(_settings.DefaultCollection) ? _settings.DefaultCollection : KnowledgeCollectionName.BotSharp;
await vectorDb.CreateCollection(collectionName, vector[0].Length);
var result = await vectorDb.Upsert(collectionName, Guid.NewGuid(), vector[0],

View file

@ -13,9 +13,9 @@ public partial class KnowledgeService
});
var db = GetVectorDb();
var textEmbedding = GetTextEmbedding();
var textEmbedding = GetTextEmbedding(collectionName);
await db.CreateCollection(collectionName, textEmbedding.Dimension);
await db.CreateCollection(collectionName, textEmbedding.GetDimension());
foreach (var line in lines)
{
var vec = await textEmbedding.GetVectorAsync(line);

View file

@ -1,5 +1,6 @@
using BotSharp.Abstraction.Graph.Models;
using BotSharp.Abstraction.VectorStorage.Models;
using System.Collections;
namespace BotSharp.Plugin.KnowledgeBase.Services;
@ -43,7 +44,7 @@ public partial class KnowledgeService
{
try
{
var textEmbedding = GetTextEmbedding();
var textEmbedding = GetTextEmbedding(collectionName);
var vector = await textEmbedding.GetVectorAsync(query);
// Vector search
@ -82,7 +83,7 @@ public partial class KnowledgeService
{
try
{
var textEmbedding = GetTextEmbedding();
var textEmbedding = GetTextEmbedding(collectionName);
var vector = await textEmbedding.GetVectorAsync(query);
var vectorDb = GetVectorDb();

View file

@ -31,13 +31,8 @@ public partial class KnowledgeService : IKnowledgeService
return db;
}
private ITextEmbedding GetTextEmbedding()
private ITextEmbedding GetTextEmbedding(string collection)
{
var embedding = _services.GetServices<ITextEmbedding>().FirstOrDefault(x => x.Provider == _settings.TextEmbedding.Provider);
if (embedding != null)
{
embedding.SetModelName(_settings.TextEmbedding.Model);
}
return embedding;
return KnowledgeSettingUtility.GetTextEmbeddingSetting(_services, collection);
}
}

View file

@ -31,4 +31,5 @@ global using BotSharp.Abstraction.Agents.Models;
global using BotSharp.Abstraction.Functions.Models;
global using BotSharp.Abstraction.Repositories;
global using BotSharp.Plugin.KnowledgeBase.Services;
global using BotSharp.Plugin.KnowledgeBase.Enum;
global using BotSharp.Plugin.KnowledgeBase.Enum;
global using BotSharp.Plugin.KnowledgeBase.Utilities;

View file

@ -0,0 +1,19 @@
namespace BotSharp.Plugin.KnowledgeBase.Utilities;
public static class KnowledgeSettingUtility
{
public static ITextEmbedding GetTextEmbeddingSetting(IServiceProvider services, string collectionName)
{
var settings = services.GetRequiredService<KnowledgeBaseSettings>();
var found = settings.Collections.FirstOrDefault(x => x.Name == collectionName)?.TextEmbedding;
if (found == null)
{
found = settings.Default.TextEmbedding;
}
var embedding = services.GetServices<ITextEmbedding>().FirstOrDefault(x => x.Provider == found.Provider);
embedding.SetModelName(found.Model);
embedding.SetDimension(found.Dimension);
return embedding;
}
}

View file

@ -7,7 +7,9 @@ public class TextEmbeddingProvider : ITextEmbedding
private LLamaEmbedder _embedder;
private readonly LlamaSharpSettings _settings;
private readonly IServiceProvider _services;
public int Dimension { get; set; } = 4096;
private const int DEFAULT_DIMENSION = 4096;
protected int _dimension = DEFAULT_DIMENSION;
public string Provider => "llama-sharp";
@ -36,4 +38,14 @@ public class TextEmbeddingProvider : ITextEmbedding
}
public void SetModelName(string model) { }
public void SetDimension(int dimension)
{
_dimension = dimension > 0 ? dimension : DEFAULT_DIMENSION;
}
public int GetDimension()
{
return _dimension;
}
}

View file

@ -14,19 +14,7 @@ public class fastTextEmbeddingProvider : ITextEmbedding
private FastTextWrapper _fastText;
private readonly IServiceProvider _services;
private int dimension;
public int Dimension
{
get
{
LoadModel();
return _fastText.GetModelDimension();
}
set
{
dimension = value;
}
}
private int _dimension;
public string Provider => "meta-ai";
@ -73,4 +61,15 @@ public class fastTextEmbeddingProvider : ITextEmbedding
}
public void SetModelName(string model) { }
public void SetDimension(int dimension)
{
LoadModel();
_dimension = _fastText.GetModelDimension();
}
public int GetDimension()
{
return _dimension;
}
}

View file

@ -10,11 +10,10 @@ public class TextEmbeddingProvider : ITextEmbedding
private const int DEFAULT_DIMENSION = 3072;
protected string _model = "text-embedding-3-large";
protected int _dimension = DEFAULT_DIMENSION;
public virtual string Provider => "openai";
public int Dimension { get; set; }
public TextEmbeddingProvider(
OpenAiSettings settings,
ILogger<TextEmbeddingProvider> logger,
@ -50,24 +49,26 @@ public class TextEmbeddingProvider : ITextEmbedding
_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 = GetDimension()
Dimensions = GetDimensionOption()
};
}
private int GetDimension()
private int GetDimensionOption()
{
var state = _services.GetRequiredService<IConversationStateService>();
var stateDimension = state.GetState("embedding_dimension");
var defaultDimension = Dimension > 0 ? Dimension : DEFAULT_DIMENSION;
if (int.TryParse(stateDimension, out var dimension))
{
return dimension > 0 ? dimension : defaultDimension;
}
return defaultDimension;
return _dimension > 0 ? _dimension : DEFAULT_DIMENSION;
}
}

View file

@ -57,12 +57,13 @@ public class IntentClassifier
return;
}
var vector = _services.GetServices<ITextEmbedding>().FirstOrDefault(x => x.Provider == _knowledgeBaseSettings.TextEmbedding.Provider);
vector.SetModelName(_knowledgeBaseSettings.TextEmbedding.Model);
var embedding = _services.GetServices<ITextEmbedding>().FirstOrDefault(x => x.Provider == _knowledgeBaseSettings.Default.TextEmbedding.Provider);
embedding.SetModelName(_knowledgeBaseSettings.Default.TextEmbedding.Model);
embedding.SetDimension(_knowledgeBaseSettings.Default.TextEmbedding.Dimension);
var layers = new List<ILayer>
{
keras.layers.InputLayer((vector.Dimension), name: "Input"),
keras.layers.InputLayer((embedding.GetDimension()), name: "Input"),
keras.layers.Dense(256, activation:"relu"),
keras.layers.Dense(256, activation:"relu"),
keras.layers.Dense(GetLabels().Length, activation: keras.activations.Softmax)
@ -136,10 +137,11 @@ public class IntentClassifier
public NDArray GetTextEmbedding(string text)
{
var knowledgeSettings = _services.GetRequiredService<KnowledgeBaseSettings>();
var embedding = _services.GetServices<ITextEmbedding>() .FirstOrDefault(x => x.Provider == knowledgeSettings.TextEmbedding.Provider);
embedding.SetModelName(knowledgeSettings.TextEmbedding.Model);
var embedding = _services.GetServices<ITextEmbedding>().FirstOrDefault(x => x.Provider == knowledgeSettings.Default.TextEmbedding.Provider);
embedding.SetModelName(knowledgeSettings.Default.TextEmbedding.Model);
embedding.SetDimension(_knowledgeBaseSettings.Default.TextEmbedding.Dimension);
var x = np.zeros((1, embedding.Dimension), dtype: np.float32);
var x = np.zeros((1, embedding.GetDimension()), dtype: np.float32);
x[0] = embedding.GetVectorAsync(text).GetAwaiter().GetResult();
return x;
}
@ -186,7 +188,7 @@ public class IntentClassifier
// Sort label to keep the same order
var uniqueLabelList = labelList.Distinct().OrderBy(x => x).ToArray();
var x = np.zeros((vectorList.Count, vector.Dimension), dtype: np.float32);
var x = np.zeros((vectorList.Count, vector.GetDimension()), dtype: np.float32);
var y = np.zeros((vectorList.Count, 1), dtype: np.float32);
for (int i = 0; i < vectorList.Count; i++)

View file

@ -24,13 +24,13 @@ namespace BotSharp.Plugin.SemanticKernel
public SemanticKernelTextEmbeddingProvider(ITextEmbeddingGenerationService embedding, IConfiguration configuration)
#pragma warning restore SKEXP0001 // Type is for evaluation purposes only and is subject to change or removal in future updates. Suppress this diagnostic to proceed.
{
this._embedding = embedding;
this._configuration = configuration;
this.Dimension = configuration.GetValue<int>("SemanticKernel:Dimension");
_embedding = embedding;
_configuration = configuration;
_dimension = configuration.GetValue<int>("SemanticKernel:Dimension");
}
/// <inheritdoc/>
public int Dimension { get; set; }
protected int _dimension;
public string Provider => "semantic-kernel";
@ -51,5 +51,15 @@ namespace BotSharp.Plugin.SemanticKernel
}
public void SetModelName(string model) { }
public void SetDimension(int dimension)
{
_dimension = dimension > 0 ? dimension : _configuration.GetValue<int>("SemanticKernel:Dimension");
}
public int GetDimension()
{
return _dimension;
}
}
}

View file

@ -264,11 +264,24 @@
"KnowledgeBase": {
"VectorDb": "Qdrant",
"GraphDb": "Default",
"DefaultCollection": "BotSharp",
"TextEmbedding": {
"Provider": "openai",
"Model": "text-embedding-3-small"
}
"Default": {
"CollectionName": "BotSharp",
"TextEmbedding": {
"Provider": "openai",
"Model": "text-embedding-3-small",
"Dimension": 1536
}
},
"Collections": [
{
"Name": "BotSharp",
"TextEmbedding": {
"Provider": "openai",
"Model": "text-embedding-3-small",
"Dimension": 1536
}
}
]
},
"SparkDesk": {