Support model_id in InstructMode

This commit is contained in:
Haiping Chen 2024-03-25 11:22:43 -05:00
parent 4d99bf788c
commit 7fb5410097
4 changed files with 50 additions and 38 deletions

View file

@ -14,6 +14,12 @@ public class MessageConfig : TruncateMessageRequest
[JsonPropertyName("model")]
public virtual string? Model { get; set; } = null;
/// <summary>
/// Model name
/// </summary>
[JsonPropertyName("model_id")]
public virtual string? ModelId { get; set; } = null;
/// <summary>
/// The sampling temperature to use that controls the apparent creativity of generated completions.
/// </summary>

View file

@ -1,4 +1,3 @@
using BotSharp.Abstraction.Agents.Models;
using BotSharp.Abstraction.MLTasks;
using BotSharp.Abstraction.MLTasks.Settings;
@ -11,31 +10,25 @@ public class CompletionProvider
string? model = null,
AgentLlmConfig? agentConfig = null)
{
var state = services.GetRequiredService<IConversationStateService>();
var agentSetting = services.GetRequiredService<AgentSettings>();
if (string.IsNullOrEmpty(provider))
{
provider = agentConfig?.Provider ?? agentSetting.LlmConfig?.Provider;
provider = state.GetState("provider", provider ?? "azure-openai");
}
if (string.IsNullOrEmpty(model))
{
model = agentConfig?.Model ?? agentSetting.LlmConfig?.Model;
model = state.GetState("model", model ?? "gpt-35-turbo-4k");
}
var settingsService = services.GetRequiredService<ILlmProviderService>();
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, agentConfig: agentConfig);
var settings = settingsService.GetSetting(provider, model);
if (settings.Type == LlmModelType.Text)
{
return GetTextCompletion(services, provider: provider, model: model);
return GetTextCompletion(services,
provider: provider,
model: model,
agentConfig: agentConfig);
}
else
{
return GetChatCompletion(services, provider: provider, model: model);
return GetChatCompletion(services,
provider: provider,
model: model,
agentConfig: agentConfig);
}
}
@ -45,20 +38,7 @@ public class CompletionProvider
AgentLlmConfig? agentConfig = null)
{
var completions = services.GetServices<IChatCompletion>();
var agentSetting = services.GetRequiredService<AgentSettings>();
var state = services.GetRequiredService<IConversationStateService>();
if (string.IsNullOrEmpty(provider))
{
provider = agentConfig?.Provider ?? agentSetting.LlmConfig?.Provider;
provider = state.GetState("provider", provider ?? "azure-openai");
}
if (string.IsNullOrEmpty(model))
{
model = agentConfig?.Model ?? agentSetting.LlmConfig?.Model;
model = state.GetState("model", model ?? "gpt-35-turbo-4k");
}
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, agentConfig: agentConfig);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)
@ -72,12 +52,11 @@ public class CompletionProvider
return completer;
}
public static ITextCompletion GetTextCompletion(IServiceProvider services,
string? provider = null,
private static (string, string) GetProviderAndModel(IServiceProvider services,
string? provider = null,
string? model = null,
AgentLlmConfig? agentConfig = null)
{
var completions = services.GetServices<ITextCompletion>();
var agentSetting = services.GetRequiredService<AgentSettings>();
var state = services.GetRequiredService<IConversationStateService>();
@ -90,9 +69,32 @@ public class CompletionProvider
if (string.IsNullOrEmpty(model))
{
model = agentConfig?.Model ?? agentSetting.LlmConfig?.Model;
model = state.GetState("model", model ?? "gpt-35-turbo-instruct");
if (state.ContainsState("model"))
{
model = state.GetState("model", model ?? "gpt-35-turbo-4k");
}
else if (state.ContainsState("model_id"))
{
var modelId = state.GetState("model_id");
var llmProviderService = services.GetRequiredService<ILlmProviderService>();
model = llmProviderService.GetProviderModel(provider, modelId)?.Name;
}
}
state.SetState("provider", provider);
state.SetState("model", model);
return (provider, model);
}
public static ITextCompletion GetTextCompletion(IServiceProvider services,
string? provider = null,
string? model = null,
AgentLlmConfig? agentConfig = null)
{
var completions = services.GetServices<ITextCompletion>();
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, agentConfig: agentConfig);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)
{

View file

@ -68,6 +68,7 @@ public partial class InstructService : IInstructService
var completer = CompletionProvider.GetCompletion(_services,
agentConfig: agent.LlmConfig);
var response = new InstructResult
{
MessageId = message.MessageId

View file

@ -25,6 +25,7 @@ public class InstructModeController : ControllerBase
input.States.ForEach(x => state.SetState(x.Split('=')[0], x.Split('=')[1]));
state.SetState("provider", input.Provider)
.SetState("model", input.Model)
.SetState("model_id", input.ModelId)
.SetState("instruction", input.Instruction)
.SetState("input_text", input.Text);
@ -45,7 +46,8 @@ public class InstructModeController : ControllerBase
var state = _services.GetRequiredService<IConversationStateService>();
input.States.ForEach(x => state.SetState(x.Split('=')[0], x.Split('=')[1]));
state.SetState("provider", input.Provider)
.SetState("model", input.Model);
.SetState("model", input.Model)
.SetState("model_id", input.ModelId);
var textCompletion = CompletionProvider.GetTextCompletion(_services);
return await textCompletion.GetCompletion(input.Text, Guid.Empty.ToString(), Guid.NewGuid().ToString());
@ -57,7 +59,8 @@ public class InstructModeController : ControllerBase
var state = _services.GetRequiredService<IConversationStateService>();
input.States.ForEach(x => state.SetState(x.Split('=')[0], x.Split('=')[1]));
state.SetState("provider", input.Provider)
.SetState("model", input.Model);
.SetState("model", input.Model)
.SetState("model_id", input.ModelId);
var textCompletion = CompletionProvider.GetChatCompletion(_services);
var message = await textCompletion.GetChatCompletions(new Agent()