Support model_id in InstructMode
This commit is contained in:
parent
4d99bf788c
commit
7fb5410097
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
{
|
||||
|
|
|
|||
|
|
@ -68,6 +68,7 @@ public partial class InstructService : IInstructService
|
|||
|
||||
var completer = CompletionProvider.GetCompletion(_services,
|
||||
agentConfig: agent.LlmConfig);
|
||||
|
||||
var response = new InstructResult
|
||||
{
|
||||
MessageId = message.MessageId
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Reference in a new issue