All user to config providers in KnowledgeBase.#77

This commit is contained in:
Haiping Chen 2023-06-26 23:08:23 -05:00
parent 884314dc2c
commit a7b263e2eb
5 changed files with 51 additions and 19 deletions

View file

@ -4,4 +4,5 @@ public class KnowledgeBaseSettings
{
public string VectorDb { get; set; }
public string TextEmbedding { get; set; }
public string TextCompletion { get; set; }
}

View file

@ -9,17 +9,14 @@ public class KnowledgeService : IKnowledgeService
{
private readonly IServiceProvider _services;
private readonly KnowledgeBaseSettings _settings;
private readonly ITextCompletion _textCompletion;
private readonly ITextChopper _textChopper;
public KnowledgeService(IServiceProvider services,
KnowledgeBaseSettings settings,
ITextCompletion textCompletion,
ITextChopper textChopper)
{
_services = services;
_settings = settings;
_textCompletion = textCompletion;
_textChopper = textChopper;
}
@ -33,14 +30,14 @@ public class KnowledgeService : IKnowledgeService
});
// Store chunks in local file system
var knowledgeStoreDir = Path.Combine("knowledge_chunks", knowledge.AgentId);
var knowledgeStoreDir = Path.Combine("knowledge_base");
if (!Directory.Exists(knowledgeStoreDir))
{
Directory.CreateDirectory(knowledgeStoreDir);
}
var knowledgePath = Path.Combine(knowledgeStoreDir, "chuncks");
File.WriteAllLines(knowledgePath + ".txt", lines);
var knowledgePath = Path.Combine(knowledgeStoreDir, knowledge.AgentId + ".txt");
File.WriteAllLines(knowledgePath, lines);
var db = GetVectorDb();
var textEmbedding = GetTextEmbedding();
@ -60,17 +57,10 @@ public class KnowledgeService : IKnowledgeService
var vector = textEmbedding.GetVector(retrievalModel.Question);
// Scan local knowledge directory
var knowledgeName = "";
var chunks = new string[0];
foreach (var file in Directory.GetFiles(Path.Combine("knowledge_chunks", retrievalModel.AgentId)))
{
knowledgeName = new FileInfo(file).Name.Split('.').First();
chunks = File.ReadAllLines(file);
}
var chunks = File.ReadAllLines(Path.Combine("knowledge_base", retrievalModel.AgentId + ".txt"));
// Vector search
var result = await GetVectorDb().Search(knowledgeName, vector);
var result = await GetVectorDb().Search(retrievalModel.AgentId, vector);
// Restore
var prompt = "";
@ -83,7 +73,7 @@ public class KnowledgeService : IKnowledgeService
prompt += "Answer the user's question based on the content provided above, and your reply should be as concise and organized as possible.\r\n";
prompt += $"Question: {retrievalModel.Question}\r\nAnswer: ";
var completion = await _textCompletion.GetCompletion(prompt);
var completion = await GetTextCompletion().GetCompletion(prompt);
return completion;
}
@ -100,4 +90,11 @@ public class KnowledgeService : IKnowledgeService
.FirstOrDefault(x => x.GetType().Name == _settings.TextEmbedding);
return embedding;
}
public ITextCompletion GetTextCompletion()
{
var textCompletion = _services.GetServices<ITextCompletion>()
.FirstOrDefault(x => x.GetType().Name == _settings.TextCompletion);
return textCompletion;
}
}

View file

@ -1,4 +1,5 @@
using BotSharp.Abstraction.VectorStorage;
using Tensorflow.NumPy;
namespace BotSharp.Core.Plugins.MemVecDb;
@ -20,7 +21,16 @@ public class MemVectorDatabase : IVectorDb
public Task<List<int>> Search(string collectionName, float[] vector, int limit = 10)
{
throw new NotImplementedException();
var cosineList = new List<double>();
for (int i = 0; i < _vectors[collectionName].Count; i++)
{
var p = CalCosineSimilarity(vector, _vectors[collectionName][i].Vector);
cosineList.Add(p);
}
var similarities = cosineList.ToArray();
var indice = np.argsort(similarities).ToArray<int>()
.Reverse().Take(limit).ToList();
return Task.FromResult(indice);
}
public Task Upsert(string collectionName, int id, float[] vector)
@ -33,4 +43,25 @@ public class MemVectorDatabase : IVectorDb
return Task.CompletedTask;
}
private double CalCosineSimilarity(float[] vector1, float[] vector2)
{
NDArray a = vector1;
NDArray b = vector2;
double num = np.dot(a, b);
if(num == 0)
{
return 0.0;
}
b = np.square(a);
var x = np.sqrt(np.sum(b));
var x3 = np.sum(np.square(vector2));
double num2 = np.sqrt(x) * np.sqrt(x3);
if(num2 == 0)
{
return 0.0;
}
return num / num2;
}
}

View file

@ -1,4 +1,4 @@
<Project Sdk="Microsoft.NET.Sdk.Web">
<Project Sdk="Microsoft.NET.Sdk.Web">
<PropertyGroup>
<TargetFramework>net6.0</TargetFramework>
@ -31,6 +31,7 @@
<ItemGroup>
<PackageReference Include="Microsoft.AspNetCore.Authentication.JwtBearer" Version="6.0.16" />
<PackageReference Include="SciSharp.TensorFlow.Redist" Version="2.11.4" />
<PackageReference Include="Swashbuckle.AspNetCore" Version="6.5.0" />
</ItemGroup>

View file

@ -68,7 +68,9 @@
},
"KnowledgeBase": {
"VectorDb": "MemVecDbProvider"
"VectorDb": "MemVectorDatabase",
"TextEmbedding": "fastTextEmbeddingProvider",
"TextCompletion": "TextCompletionProvider"
},
"PluginLoader": {