@@ -2,10 +2,12 @@ package openai
22
33import (
44 "context"
5+ "encoding/json"
56 "fmt"
67
78 "github.com/mudler/LocalAI/core/backend"
89 "github.com/mudler/LocalAI/core/config"
10+ "github.com/mudler/LocalAI/core/http/endpoints/openai/types"
911 "github.com/mudler/LocalAI/core/schema"
1012 "github.com/mudler/LocalAI/core/templates"
1113 "github.com/mudler/LocalAI/pkg/functions"
@@ -63,14 +65,18 @@ func (m *transcriptOnlyModel) Transcribe(ctx context.Context, audio, language st
6365 return backend .ModelTranscription (audio , language , translate , diarize , prompt , m .modelLoader , * m .TranscriptionConfig , m .appConfig )
6466}
6567
66- func (m * transcriptOnlyModel ) Predict (ctx context.Context , messages schema.Messages , images , videos , audios []string , tokenCallback func (string , backend.TokenUsage ) bool , tools string , toolChoice string , logprobs * int , topLogprobs * int , logitBias map [string ]float64 ) (func () (backend.LLMResponse , error ), error ) {
68+ func (m * transcriptOnlyModel ) Predict (ctx context.Context , messages schema.Messages , images , videos , audios []string , tokenCallback func (string , backend.TokenUsage ) bool , tools []types. ToolUnion , toolChoice * types. ToolChoiceUnion , logprobs * int , topLogprobs * int , logitBias map [string ]float64 ) (func () (backend.LLMResponse , error ), error ) {
6769 return nil , fmt .Errorf ("predict operation not supported in transcript-only mode" )
6870}
6971
7072func (m * transcriptOnlyModel ) TTS (ctx context.Context , text , voice , language string ) (string , * proto.Result , error ) {
7173 return "" , nil , fmt .Errorf ("TTS not supported in transcript-only mode" )
7274}
7375
76+ func (m * transcriptOnlyModel ) PredictConfig () * config.ModelConfig {
77+ return nil
78+ }
79+
7480func (m * wrappedModel ) VAD (ctx context.Context , request * schema.VADRequest ) (* schema.VADResponse , error ) {
7581 return backend .VAD (request , ctx , m .modelLoader , m .appConfig , * m .VADConfig )
7682}
@@ -79,28 +85,116 @@ func (m *wrappedModel) Transcribe(ctx context.Context, audio, language string, t
7985 return backend .ModelTranscription (audio , language , translate , diarize , prompt , m .modelLoader , * m .TranscriptionConfig , m .appConfig )
8086}
8187
82- func (m * wrappedModel ) Predict (ctx context.Context , messages schema.Messages , images , videos , audios []string , tokenCallback func (string , backend.TokenUsage ) bool , tools string , toolChoice string , logprobs * int , topLogprobs * int , logitBias map [string ]float64 ) (func () (backend.LLMResponse , error ), error ) {
88+ func (m * wrappedModel ) Predict (ctx context.Context , messages schema.Messages , images , videos , audios []string , tokenCallback func (string , backend.TokenUsage ) bool , tools []types. ToolUnion , toolChoice * types. ToolChoiceUnion , logprobs * int , topLogprobs * int , logitBias map [string ]float64 ) (func () (backend.LLMResponse , error ), error ) {
8389 input := schema.OpenAIRequest {
8490 Messages : messages ,
8591 }
8692
8793 var predInput string
94+ var funcs []functions.Function
8895 if ! m .LLMConfig .TemplateConfig .UseTokenizerTemplate {
89- predInput = m .evaluator .TemplateMessages (input , input .Messages , m .LLMConfig , []functions.Function {}, false )
96+ if len (tools ) > 0 {
97+ for _ , t := range tools {
98+ if t .Function != nil {
99+ var params map [string ]any
100+
101+ switch p := t .Function .Parameters .(type ) {
102+ case map [string ]any :
103+ params = p
104+ case string :
105+ if err := json .Unmarshal ([]byte (p ), & params ); err != nil {
106+ xlog .Warn ("Failed to parse parameters JSON string" , "error" , err , "function" , t .Function .Name )
107+ }
108+ }
109+
110+ funcs = append (funcs , functions.Function {
111+ Name : t .Function .Name ,
112+ Description : t .Function .Description ,
113+ Parameters : params ,
114+ })
115+ }
116+ }
117+ }
118+
119+ predInput = m .evaluator .TemplateMessages (input , input .Messages , m .LLMConfig , funcs , len (funcs ) > 0 )
90120
91121 xlog .Debug ("Prompt (after templating)" , "prompt" , predInput )
92122 if m .LLMConfig .Grammar != "" {
93123 xlog .Debug ("Grammar" , "grammar" , m .LLMConfig .Grammar )
94124 }
95125 }
96126
97- return backend .ModelInference (ctx , predInput , messages , images , videos , audios , m .modelLoader , m .LLMConfig , m .confLoader , m .appConfig , tokenCallback , tools , toolChoice , logprobs , topLogprobs , logitBias , )
127+ // Generate grammar for function calling if tools are provided and grammar generation is enabled
128+ shouldUseFn := len (tools ) > 0 && m .LLMConfig .ShouldUseFunctions ()
129+
130+ if ! m .LLMConfig .FunctionsConfig .GrammarConfig .NoGrammar && shouldUseFn {
131+ // Allow the user to set custom actions via config file
132+ noActionName := "answer"
133+ noActionDescription := "use this action to answer without performing any action"
134+
135+ if m .LLMConfig .FunctionsConfig .NoActionFunctionName != "" {
136+ noActionName = m .LLMConfig .FunctionsConfig .NoActionFunctionName
137+ }
138+ if m .LLMConfig .FunctionsConfig .NoActionDescriptionName != "" {
139+ noActionDescription = m .LLMConfig .FunctionsConfig .NoActionDescriptionName
140+ }
141+
142+ noActionGrammar := functions.Function {
143+ Name : noActionName ,
144+ Description : noActionDescription ,
145+ Parameters : map [string ]interface {}{
146+ "properties" : map [string ]interface {}{
147+ "message" : map [string ]interface {}{
148+ "type" : "string" ,
149+ "description" : "The message to reply the user with" ,
150+ },
151+ },
152+ },
153+ }
154+
155+ if ! m .LLMConfig .FunctionsConfig .DisableNoAction {
156+ funcs = append (funcs , noActionGrammar )
157+ }
158+
159+ // Force picking one of the functions by the request
160+ if m .LLMConfig .FunctionToCall () != "" {
161+ funcs = functions .Functions (funcs ).Select (m .LLMConfig .FunctionToCall ())
162+ }
163+
164+ // Generate grammar from function definitions
165+ jsStruct := functions .Functions (funcs ).ToJSONStructure (m .LLMConfig .FunctionsConfig .FunctionNameKey , m .LLMConfig .FunctionsConfig .FunctionNameKey )
166+ g , err := jsStruct .Grammar (m .LLMConfig .FunctionsConfig .GrammarOptions ()... )
167+ if err == nil {
168+ m .LLMConfig .Grammar = g
169+ xlog .Debug ("Generated grammar for function calling" , "grammar" , g )
170+ } else {
171+ xlog .Error ("Failed generating grammar" , "error" , err )
172+ }
173+ }
174+
175+ var toolsJSON string
176+ if len (tools ) > 0 {
177+ b , _ := json .Marshal (tools )
178+ toolsJSON = string (b )
179+ }
180+
181+ var toolChoiceJSON string
182+ if toolChoice != nil {
183+ b , _ := json .Marshal (toolChoice )
184+ toolChoiceJSON = string (b )
185+ }
186+
187+ return backend .ModelInference (ctx , predInput , messages , images , videos , audios , m .modelLoader , m .LLMConfig , m .confLoader , m .appConfig , tokenCallback , toolsJSON , toolChoiceJSON , logprobs , topLogprobs , logitBias , )
98188}
99189
100190func (m * wrappedModel ) TTS (ctx context.Context , text , voice , language string ) (string , * proto.Result , error ) {
101191 return backend .ModelTTS (text , voice , language , m .modelLoader , m .appConfig , * m .TTSConfig )
102192}
103193
194+ func (m * wrappedModel ) PredictConfig () * config.ModelConfig {
195+ return m .LLMConfig
196+ }
197+
104198func newTranscriptionOnlyModel (pipeline * config.Pipeline , cl * config.ModelConfigLoader , ml * model.ModelLoader , appConfig * config.ApplicationConfig ) (Model , * config.ModelConfig , error ) {
105199 cfgVAD , err := cl .LoadModelConfigFileByName (pipeline .VAD , ml .ModelPath )
106200 if err != nil {
0 commit comments