|
|
|
@ -33,7 +33,7 @@ type LlamaServer interface { |
|
|
|
Ping(ctx context.Context) error |
|
|
|
WaitUntilRunning(ctx context.Context) error |
|
|
|
Completion(ctx context.Context, req CompletionRequest, fn func(CompletionResponse)) error |
|
|
|
Embed(ctx context.Context, input []string) ([][]float32, error) |
|
|
|
Embed(ctx context.Context, input []string) (*EmbedResponse, error) |
|
|
|
Tokenize(ctx context.Context, content string) ([]int, error) |
|
|
|
Detokenize(ctx context.Context, tokens []int) (string, error) |
|
|
|
Close() error |
|
|
|
@ -879,10 +879,11 @@ type EmbedRequest struct { |
|
|
|
} |
|
|
|
|
|
|
|
type EmbedResponse struct { |
|
|
|
Embedding [][]float32 `json:"embedding"` |
|
|
|
Embedding [][]float32 `json:"embedding"` |
|
|
|
PromptEvalCount int `json:"prompt_n"` |
|
|
|
} |
|
|
|
|
|
|
|
func (s *llmServer) Embed(ctx context.Context, input []string) ([][]float32, error) { |
|
|
|
func (s *llmServer) Embed(ctx context.Context, input []string) (*EmbedResponse, error) { |
|
|
|
if err := s.sem.Acquire(ctx, 1); err != nil { |
|
|
|
slog.Error("Failed to acquire semaphore", "error", err) |
|
|
|
return nil, err |
|
|
|
@ -924,12 +925,12 @@ func (s *llmServer) Embed(ctx context.Context, input []string) ([][]float32, err |
|
|
|
return nil, fmt.Errorf("%s", body) |
|
|
|
} |
|
|
|
|
|
|
|
var embedding EmbedResponse |
|
|
|
if err := json.Unmarshal(body, &embedding); err != nil { |
|
|
|
var e EmbedResponse |
|
|
|
if err := json.Unmarshal(body, &e); err != nil { |
|
|
|
return nil, fmt.Errorf("unmarshal tokenize response: %w", err) |
|
|
|
} |
|
|
|
|
|
|
|
return embedding.Embedding, nil |
|
|
|
return &e, nil |
|
|
|
} |
|
|
|
|
|
|
|
type TokenizeRequest struct { |
|
|
|
|