This commit is contained in:
Haiping Chen 2023-10-11 19:59:09 -05:00
parent 7fbff8c777
commit fbd3c6f0ca
5 changed files with 63 additions and 18 deletions

View file

@ -9,10 +9,12 @@ public class RoutingArgs
public string Reason { get; set; } = string.Empty; public string Reason { get; set; } = string.Empty;
[JsonPropertyName("answer")] [JsonPropertyName("answer")]
public string Answer { get; set; } = string.Empty; [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
public string? Answer { get; set; }
[JsonPropertyName("agent")] [JsonPropertyName("agent")]
public string AgentName { get; set; } = string.Empty; [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
public string? AgentName { get; set; }
public override string ToString() public override string ToString()
{ {

View file

@ -9,13 +9,23 @@ public partial class RoutingService
{ {
public async Task<FunctionCallFromLlm> GetNextInstruction(string prompt) public async Task<FunctionCallFromLlm> GetNextInstruction(string prompt)
{ {
var responseFormat = _settings.EnableReasoning ? var responseRouteToAgent = _settings.EnableReasoning ?
JsonSerializer.Serialize(new FunctionCallFromLlm()) : JsonSerializer.Serialize(new FunctionCallFromLlm()) :
JsonSerializer.Serialize(new RoutingArgs JsonSerializer.Serialize(new RoutingArgs
{ {
Function = "route_to_agent" Function = "route_to_agent",
AgentName = ""
}); });
var content = $"{prompt} Response must be in JSON format {responseFormat}";
var responseResponseToUser = _settings.EnableReasoning ?
JsonSerializer.Serialize(new FunctionCallFromLlm()) :
JsonSerializer.Serialize(new RoutingArgs
{
Function = "response_to_user",
Answer = ""
});
var content = $"{prompt}\r\nIf need to route to specific agent, output in JSON {responseRouteToAgent}.\r\nIf you can handle the request directly, output in JSON {responseResponseToUser}.";
var state = _services.GetRequiredService<IConversationStateService>(); var state = _services.GetRequiredService<IConversationStateService>();
@ -29,8 +39,34 @@ public partial class RoutingService
model: _settings.Model); model: _settings.Model);
content = _routerInstance.Router.Instruction + "\r\n\r\n" + content + "\r\nResponse: "; content = _routerInstance.Router.Instruction + "\r\n\r\n" + content + "\r\nResponse: ";
var text = await completion.GetCompletion(content);
response = new RoleDialogModel(AgentRole.Assistant, text); int retryCount = 0;
while (retryCount < 3)
{
try
{
var text = await completion.GetCompletion(content);
response = new RoleDialogModel(AgentRole.Assistant, text);
var pattern = @"\{(?:[^{}]|(?<open>\{)|(?<-open>\}))+(?(open)(?!))\}";
response.Content = Regex.Match(response.Content, pattern).Value;
args = JsonSerializer.Deserialize<FunctionCallFromLlm>(response.Content);
break;
}
catch (Exception ex)
{
_logger.LogError($"{ex.Message}: {response.Content}");
args.Function = "response_to_user";
args.Answer = ex.Message;
args.AgentName = _settings.RouterName;
content += "\r\nPlease response in JSON format.";
}
finally
{
retryCount++;
}
}
} }
else else
{ {

View file

@ -77,7 +77,6 @@ public partial class RoutingService : IRoutingService
loopCount++; loopCount++;
var prompt = _settings.EnableReasoning ? "Tell me the next step?" : "Which agent is suitable to handle user's request based on the CONVERSATION?"; var prompt = _settings.EnableReasoning ? "Tell me the next step?" : "Which agent is suitable to handle user's request based on the CONVERSATION?";
prompt += " Or you can handle without asking specific agent.";
var inst = await GetNextInstruction(prompt); var inst = await GetNextInstruction(prompt);
inst.Question = inst.Question ?? message; inst.Question = inst.Question ?? message;

View file

@ -23,15 +23,18 @@ public class ChatCompletionProvider : IChatCompletion
private readonly IServiceProvider _services; private readonly IServiceProvider _services;
private readonly ILogger _logger; private readonly ILogger _logger;
private readonly LlamaSharpSettings _settings; private readonly LlamaSharpSettings _settings;
private readonly ITokenStatistics _tokenStatistics;
private string _model; private string _model;
public ChatCompletionProvider(IServiceProvider services, public ChatCompletionProvider(IServiceProvider services,
ILogger<ChatCompletionProvider> logger, ILogger<ChatCompletionProvider> logger,
LlamaSharpSettings settings) LlamaSharpSettings settings,
ITokenStatistics tokenStatistics)
{ {
_services = services; _services = services;
_logger = logger; _logger = logger;
_settings = settings; _settings = settings;
_tokenStatistics = tokenStatistics;
} }
public string Provider => "llama-sharp"; public string Provider => "llama-sharp";
@ -41,11 +44,8 @@ public class ChatCompletionProvider : IChatCompletion
var content = string.Join("\r\n", conversations.Select(x => $"{x.Role}: {x.Content}")).Trim(); var content = string.Join("\r\n", conversations.Select(x => $"{x.Role}: {x.Content}")).Trim();
content += $"\r\n{AgentRole.Assistant}: "; content += $"\r\n{AgentRole.Assistant}: ";
var state = _services.GetRequiredService<IConversationStateService>();
var model = state.GetState("model", _settings.DefaultModel);
var llama = _services.GetRequiredService<LlamaAiModel>(); var llama = _services.GetRequiredService<LlamaAiModel>();
llama.LoadModel(model); llama.LoadModel(_model);
var executor = llama.GetStatelessExecutor(); var executor = llama.GetStatelessExecutor();
var inferenceParams = new InferenceParams() var inferenceParams = new InferenceParams()
@ -65,11 +65,13 @@ public class ChatCompletionProvider : IChatCompletion
_logger.LogInformation(prompt); _logger.LogInformation(prompt);
} }
_tokenStatistics.StartTimer();
foreach (var response in executor.Infer(prompt, inferenceParams)) foreach (var response in executor.Infer(prompt, inferenceParams))
{ {
Console.Write(response); Console.Write(response);
totalResponse += response; totalResponse += response;
} }
_tokenStatistics.StopTimer();
foreach (var anti in inferenceParams.AntiPrompts) foreach (var anti in inferenceParams.AntiPrompts)
{ {

View file

@ -5,6 +5,7 @@ using BotSharp.Plugins.LLamaSharp;
using LLama; using LLama;
using LLama.Common; using LLama.Common;
using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.DependencyInjection;
using Microsoft.Extensions.Logging;
using System; using System;
using System.Threading.Tasks; using System.Threading.Tasks;
@ -13,34 +14,39 @@ namespace BotSharp.Plugin.LLamaSharp.Providers;
public class TextCompletionProvider : ITextCompletion public class TextCompletionProvider : ITextCompletion
{ {
private readonly IServiceProvider _services; private readonly IServiceProvider _services;
private readonly ILogger _logger;
private readonly LlamaSharpSettings _settings; private readonly LlamaSharpSettings _settings;
private readonly ITokenStatistics _tokenStatistics;
private string _model; private string _model;
public string Provider => "llama-sharp"; public string Provider => "llama-sharp";
public TextCompletionProvider(IServiceProvider services, public TextCompletionProvider(IServiceProvider services,
LlamaSharpSettings settings) ILogger<TextCompletionProvider> logger,
LlamaSharpSettings settings,
ITokenStatistics tokenStatistics)
{ {
_services = services; _services = services;
_logger = logger;
_settings = settings; _settings = settings;
_tokenStatistics = tokenStatistics;
} }
public Task<string> GetCompletion(string text) public Task<string> GetCompletion(string text)
{ {
var state = _services.GetRequiredService<IConversationStateService>();
var model = state.GetState("model", _settings.DefaultModel);
var llama = _services.GetRequiredService<LlamaAiModel>(); var llama = _services.GetRequiredService<LlamaAiModel>();
llama.LoadModel(model); llama.LoadModel(_model);
var executor = new InstructExecutor(llama.Model.CreateContext(llama.Params)); var executor = new InstructExecutor(llama.Model.CreateContext(llama.Params));
var inferenceParams = new InferenceParams() { Temperature = 0.5f, MaxTokens = 128 }; var inferenceParams = new InferenceParams() { Temperature = 0.5f, MaxTokens = 128 };
_tokenStatistics.StartTimer();
string totalResponse = ""; string totalResponse = "";
foreach (var response in executor.Infer(text, inferenceParams)) foreach (var response in executor.Infer(text, inferenceParams))
{ {
Console.Write(response); Console.Write(response);
totalResponse += response; totalResponse += response;
} }
_tokenStatistics.StopTimer();
return Task.FromResult(totalResponse); return Task.FromResult(totalResponse);
} }