Fix LLamaSharpPlugin.TextEmbeddingProvider
This commit is contained in:
parent
ba1ff71020
commit
e81e01f8e3
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -285,3 +285,4 @@ __pycache__/
|
|||
*.xsd.cs
|
||||
data
|
||||
/docs/_build
|
||||
*.bin
|
||||
|
|
|
|||
|
|
@ -23,10 +23,11 @@
|
|||
|
||||
<ItemGroup>
|
||||
<PackageReference Include="Microsoft.AspNetCore.Http.Abstractions" Version="2.2.0" />
|
||||
<PackageReference Include="Microsoft.Extensions.Configuration.Binder" Version="6.0.0" />
|
||||
<PackageReference Include="Microsoft.Extensions.DependencyInjection.Abstractions" Version="6.0.0" />
|
||||
<PackageReference Include="Microsoft.Extensions.Configuration.Binder" Version="7.0.4" />
|
||||
<PackageReference Include="Microsoft.Extensions.DependencyInjection.Abstractions" Version="7.0.0" />
|
||||
<PackageReference Include="Microsoft.Extensions.Logging.Abstractions" Version="7.0.1" />
|
||||
<PackageReference Include="System.ComponentModel.Annotations" Version="5.0.0" />
|
||||
<PackageReference Include="System.Text.Json" Version="6.0.0" />
|
||||
<PackageReference Include="System.Text.Json" Version="7.0.3" />
|
||||
</ItemGroup>
|
||||
|
||||
</Project>
|
||||
|
|
|
|||
|
|
@ -1,10 +1,16 @@
|
|||
using System.Text.Json.Serialization;
|
||||
|
||||
namespace BotSharp.Abstraction.Users.Models;
|
||||
|
||||
public class Token
|
||||
{
|
||||
[JsonPropertyName("access_token")]
|
||||
public string AccessToken { get; set; } = string.Empty;
|
||||
[JsonPropertyName("refresh_token")]
|
||||
public string RefreshToken { get; set; } = string.Empty;
|
||||
[JsonPropertyName("token_type")]
|
||||
public string TokenType { get; set; } = string.Empty;
|
||||
[JsonPropertyName("expires")]
|
||||
public int ExpireTime { get; set; }
|
||||
public string Scope { get; set; } = string.Empty;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -75,8 +75,7 @@
|
|||
<ItemGroup>
|
||||
<PackageReference Include="Colorful.Console" Version="1.2.15" />
|
||||
<PackageReference Include="EntityFrameworkCore.BootKit" Version="6.2.1" />
|
||||
<PackageReference Include="LLamaSharp" Version="0.4.0" />
|
||||
<PackageReference Include="LLamaSharp.Backend.Cuda11" Version="0.3.0" />
|
||||
<PackageReference Include="LLamaSharp" Version="0.4.2-preview" />
|
||||
<PackageReference Include="PdfPig" Version="0.1.8" />
|
||||
<PackageReference Include="TensorFlow.Keras" Version="0.11.2" />
|
||||
<PackageReference Include="Microsoft.AspNetCore.Mvc.Core" Version="2.2.5" />
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ public class KnowledgeController : ControllerBase, IApiAdapter
|
|||
}
|
||||
|
||||
[HttpPost("/knowledge/{agentId}")]
|
||||
public async Task<IActionResult> FeedKnowledge([FromRoute] string agentId, List<IFormFile> files)
|
||||
public async Task<IActionResult> FeedKnowledge([FromRoute] string agentId, List<IFormFile> files, [FromQuery] int? startPageNum, [FromQuery] int? endPageNum)
|
||||
{
|
||||
long size = files.Sum(f => f.Length);
|
||||
|
||||
|
|
@ -52,6 +52,16 @@ public class KnowledgeController : ControllerBase, IApiAdapter
|
|||
var content = "";
|
||||
foreach (Page page in document.GetPages())
|
||||
{
|
||||
if (startPageNum.HasValue && page.Number < startPageNum.Value)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
if (endPageNum.HasValue && page.Number > endPageNum.Value)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
content += page.Text;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -24,8 +24,8 @@ public class KnowledgeService : IKnowledgeService
|
|||
var idStart = 0;
|
||||
var lines = _textChopper.Chop(knowledge.Content, new ChunkOption
|
||||
{
|
||||
Size = 256,
|
||||
Conjunction = 5,
|
||||
Size = 1024,
|
||||
Conjunction = 32,
|
||||
SplitByWord = true,
|
||||
});
|
||||
|
||||
|
|
@ -38,6 +38,7 @@ public class KnowledgeService : IKnowledgeService
|
|||
var vec = textEmbedding.GetVector(line);
|
||||
await db.Upsert(knowledge.AgentId, idStart, vec, line);
|
||||
idStart++;
|
||||
Console.WriteLine($"Saved vector {idStart}/{lines.Count}: {line}\n");
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -50,7 +51,7 @@ public class KnowledgeService : IKnowledgeService
|
|||
var result = await GetVectorDb().Search(retrievalModel.AgentId, vector, limit: 10);
|
||||
|
||||
// Restore
|
||||
return "### Helpful domain knowledges:\r\n" + string.Join("\n", result.Select((x, i) => $"{i + 1}: {x}"));
|
||||
return string.Join("\n\n", result.Select((x, i) => $"{i + 1}: {x.Trim()}"));
|
||||
}
|
||||
|
||||
public async Task<string> GetAnswer(KnowledgeRetrievalModel retrievalModel)
|
||||
|
|
@ -58,8 +59,13 @@ public class KnowledgeService : IKnowledgeService
|
|||
// Restore
|
||||
var prompt = await GetKnowledges(retrievalModel);
|
||||
|
||||
prompt += "\r\n### Answer user's question by utilizing the helpful domain knowledges above.\r\n";
|
||||
prompt += $"\r\nQuestion: {retrievalModel.Question}\r\nAnswer: ";
|
||||
var sb = new StringBuilder(prompt);
|
||||
sb.AppendLine();
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("### Answer question based on the given information above. Try to response in bullet points if necessary. Please keep your answers concise and free of irrelevant information.");
|
||||
sb.AppendLine($"Question: {retrievalModel.Question}");
|
||||
sb.AppendLine("Answer: ");
|
||||
prompt = sb.ToString().Trim();
|
||||
|
||||
var completion = await GetTextCompletion().GetCompletion(prompt);
|
||||
return completion;
|
||||
|
|
@ -68,14 +74,14 @@ public class KnowledgeService : IKnowledgeService
|
|||
public IVectorDb GetVectorDb()
|
||||
{
|
||||
var db = _services.GetServices<IVectorDb>()
|
||||
.FirstOrDefault(x => x.GetType().Name == _settings.VectorDb);
|
||||
.FirstOrDefault(x => x.GetType().FullName.EndsWith(_settings.VectorDb));
|
||||
return db;
|
||||
}
|
||||
|
||||
public ITextEmbedding GetTextEmbedding()
|
||||
{
|
||||
var embedding = _services.GetServices<ITextEmbedding>()
|
||||
.FirstOrDefault(x => x.GetType().Name == _settings.TextEmbedding);
|
||||
.FirstOrDefault(x => x.GetType().FullName.EndsWith(_settings.TextEmbedding));
|
||||
return embedding;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ public class LLamaSharpPlugin : IBotSharpPlugin
|
|||
services.AddSingleton(x => llamaSharpSettings);
|
||||
|
||||
services.AddSingleton<LlamaAiModel>();
|
||||
services.AddScoped<ITextEmbedding, TextEmbeddingProvider>();
|
||||
services.AddSingleton<ITextEmbedding, TextEmbeddingProvider>();
|
||||
services.AddScoped<ITextCompletion, TextCompletionProvider>();
|
||||
services.AddScoped<IChatCompletion, ChatCompletionProvider>();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,20 +6,24 @@ namespace BotSharp.Core.Plugins.LLamaSharp;
|
|||
|
||||
public class TextEmbeddingProvider : ITextEmbedding
|
||||
{
|
||||
private LLamaEmbedder _embedder;
|
||||
private readonly LlamaSharpSettings _settings;
|
||||
private readonly IServiceProvider _services;
|
||||
public int Dimension => throw new NotImplementedException();
|
||||
public int Dimension => 4096;
|
||||
|
||||
public TextEmbeddingProvider(IServiceProvider services)
|
||||
public TextEmbeddingProvider(IServiceProvider services, LlamaSharpSettings settings)
|
||||
{
|
||||
_services = services;
|
||||
_settings = settings;
|
||||
}
|
||||
|
||||
public float[] GetVector(string text)
|
||||
{
|
||||
var llama = _services.GetRequiredService<LlamaAiModel>();
|
||||
if (_embedder == null)
|
||||
{
|
||||
_embedder = new LLamaEmbedder(new ModelParams(_settings.ModelPath));
|
||||
}
|
||||
|
||||
var executor = new LLamaEmbedder(new ModelParams(llama.Settings.ModelPath));
|
||||
|
||||
return executor.GetEmbeddings(text);
|
||||
return _embedder.GetEmbeddings(text);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,4 @@
|
|||
using BotSharp.Abstraction.VectorStorage;
|
||||
using System.Collections;
|
||||
using System.IO;
|
||||
using System.Numerics;
|
||||
using Tensorflow;
|
||||
using Tensorflow.NumPy;
|
||||
|
||||
namespace BotSharp.Core.Plugins.MemVecDb;
|
||||
|
|
@ -67,12 +63,15 @@ public class MemVectorDatabase : IVectorDb
|
|||
private float[] CalCosineSimilarity(float[] vec, List<VecRecord> records)
|
||||
{
|
||||
var similarities = new float[records.Count];
|
||||
var a = vec;
|
||||
var normA = np.linalg.norm(a);
|
||||
|
||||
for (int i = 0; i < records.Count; i++)
|
||||
{
|
||||
var a = vec;
|
||||
var b = records[i].Vector;
|
||||
similarities[i] = np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b));
|
||||
similarities[i] = np.dot(a, b) / (normA * np.linalg.norm(b));
|
||||
}
|
||||
|
||||
return similarities;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,17 +4,20 @@ using BotSharp.Abstraction.MLTasks;
|
|||
using System;
|
||||
using System.Threading.Tasks;
|
||||
using BotSharp.Plugin.AzureOpenAI.Settings;
|
||||
using Microsoft.Extensions.Logging;
|
||||
|
||||
namespace BotSharp.Plugin.AzureOpenAI.Providers;
|
||||
|
||||
public class TextCompletionProvider : ITextCompletion
|
||||
{
|
||||
private readonly AzureOpenAiSettings _settings;
|
||||
private readonly ILogger _logger;
|
||||
bool _useAzureOpenAI = true;
|
||||
|
||||
public TextCompletionProvider(AzureOpenAiSettings settings)
|
||||
public TextCompletionProvider(AzureOpenAiSettings settings, ILogger<TextCompletionProvider> logger)
|
||||
{
|
||||
_settings = settings;
|
||||
_logger = logger;
|
||||
}
|
||||
|
||||
public async Task<string> GetCompletion(string text)
|
||||
|
|
@ -26,8 +29,8 @@ public class TextCompletionProvider : ITextCompletion
|
|||
{
|
||||
text
|
||||
},
|
||||
Temperature = 0.5f,
|
||||
MaxTokens = 128
|
||||
Temperature = 1f,
|
||||
MaxTokens = 256
|
||||
};
|
||||
|
||||
var response = await client.GetCompletionsAsync(
|
||||
|
|
@ -41,6 +44,8 @@ public class TextCompletionProvider : ITextCompletion
|
|||
completion += t.Text;
|
||||
};
|
||||
|
||||
_logger.LogInformation(text + completion);
|
||||
|
||||
return completion.Trim();
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@
|
|||
|
||||
<ItemGroup>
|
||||
<PackageReference Include="Microsoft.AspNetCore.Mvc.Core" Version="2.2.5" />
|
||||
<PackageReference Include="System.Text.Json" Version="7.0.2" />
|
||||
</ItemGroup>
|
||||
|
||||
<ItemGroup>
|
||||
|
|
|
|||
|
|
@ -31,11 +31,13 @@ public class WebhookController : ControllerBase
|
|||
_services = services;
|
||||
}
|
||||
|
||||
[HttpGet("/webhook")]
|
||||
[HttpGet("/messenger/webhook/{agentId}")]
|
||||
public string Verificate([FromQuery(Name = "hub.mode")] string mode,
|
||||
[FromQuery(Name = "hub.verify_token")] string token,
|
||||
[FromQuery(Name = "hub.challenge")] string challenge)
|
||||
[FromQuery(Name = "hub.challenge")] string challenge,
|
||||
[FromRoute] string agentId)
|
||||
{
|
||||
Console.WriteLine(agentId);
|
||||
return challenge;
|
||||
}
|
||||
|
||||
|
|
@ -43,11 +45,12 @@ public class WebhookController : ControllerBase
|
|||
/// https://developers.facebook.com/docs/messenger-platform/webhooks
|
||||
/// </summary>
|
||||
/// <returns></returns>
|
||||
[HttpPost("/webhook")]
|
||||
public async Task<ActionResult<WebhookResponse>> Messages()
|
||||
[HttpPost("/messenger/webhook/{agentId}")]
|
||||
public async Task<ActionResult<WebhookResponse>> Messages([FromRoute] string agentId)
|
||||
{
|
||||
using var stream = new StreamReader(Request.Body);
|
||||
var body = await stream.ReadToEndAsync();
|
||||
Console.WriteLine(body);
|
||||
var req = JsonSerializer.Deserialize<WebhookRequest>(body, new JsonSerializerOptions
|
||||
{
|
||||
PropertyNameCaseInsensitive = true,
|
||||
|
|
@ -66,7 +69,7 @@ public class WebhookController : ControllerBase
|
|||
string content = "";
|
||||
var sessionId = req.Entry[0].Messaging[0].Sender.Id;
|
||||
var input = req.Entry[0].Messaging[0].Message.Text;
|
||||
var result = await conv.SendMessage("", sessionId, new RoleDialogModel("user", input), async msg =>
|
||||
var result = await conv.SendMessage(agentId, sessionId, new RoleDialogModel("user", input), async msg =>
|
||||
{
|
||||
content = msg.Content;
|
||||
});
|
||||
|
|
|
|||
|
|
@ -62,10 +62,10 @@ public class QdrantDb : IVectorDb
|
|||
public async Task Upsert(string collectionName, int id, float[] vector, string text)
|
||||
{
|
||||
// Insert vectors
|
||||
/*await _client.Upsert(collectionName, points: new List<PointStruct>
|
||||
await _client.Upsert(collectionName, points: new List<PointStruct>
|
||||
{
|
||||
new PointStruct(id: id, vector: vector)
|
||||
});*/
|
||||
});
|
||||
|
||||
// Store chunks in local file system
|
||||
var agentService = _services.GetRequiredService<IAgentService>();
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@
|
|||
</ItemGroup>
|
||||
|
||||
<ItemGroup>
|
||||
<PackageReference Include="LLamaSharp.Backend.Cuda11" Version="0.4.2-preview" />
|
||||
<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" />
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@
|
|||
|
||||
"LlamaSharp": {
|
||||
"Interactive": true,
|
||||
"ModelPath": "C:\\Users\\haipi\\Downloads\\wizard-vicuna-13B.ggmlv3.q8_0.bin",
|
||||
"ModelPath": "C:/Users/haipi/Downloads/llama-2-7b-chat.ggmlv3.q3_K_S.bin",
|
||||
"MaxContextLength": 1024,
|
||||
"NumberOfGpuLayer": 10
|
||||
},
|
||||
|
|
@ -44,6 +44,13 @@
|
|||
}
|
||||
},
|
||||
|
||||
"MetaMessenger": {
|
||||
"Endpoint": "https://graph.facebook.com",
|
||||
"ApiVersion": "v17.0",
|
||||
"PageId": "",
|
||||
"PageAccessToken": ""
|
||||
},
|
||||
|
||||
"Database": {
|
||||
"MongoDb": {
|
||||
"Master": "mongodb://localhost:27017/chat-ui"
|
||||
|
|
@ -71,7 +78,9 @@
|
|||
|
||||
"KnowledgeBase": {
|
||||
"VectorDb": "MemVectorDatabase",
|
||||
// "VectorDb": "QdrantDb",
|
||||
"TextEmbedding": "fastTextEmbeddingProvider",
|
||||
// "TextEmbedding": "LLamaSharp.TextEmbeddingProvider",
|
||||
"TextCompletion": "AzureOpenAI.Providers.TextCompletionProvider"
|
||||
// "TextCompletion": "LLamaSharp.TextCompletionProvider"
|
||||
},
|
||||
|
|
|
|||
BIN
tests/Dishwasher-Whirlpool.pdf
Normal file
BIN
tests/Dishwasher-Whirlpool.pdf
Normal file
Binary file not shown.
Loading…
Reference in a new issue