Fix LLamaSharpPlugin.TextEmbeddingProvider

This commit is contained in:
Haiping Chen 2023-08-07 05:21:31 -05:00
parent ba1ff71020
commit e81e01f8e3
16 changed files with 81 additions and 38 deletions

1
.gitignore vendored
View file

@ -285,3 +285,4 @@ __pycache__/
*.xsd.cs
data
/docs/_build
*.bin

View file

@ -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>

View file

@ -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;
}

View file

@ -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" />

View file

@ -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;
}

View file

@ -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;
}

View file

@ -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>();
}

View file

@ -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);
}
}

View file

@ -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;
}
}

View file

@ -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();
}

View file

@ -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>

View file

@ -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;
});

View file

@ -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>();

View file

@ -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" />

View file

@ -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"
},

Binary file not shown.