Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 4 additions & 21 deletions examples/kv_events/online/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@
// ChatCompletionsRequest holds the fields needed for chat-completions rendering.
type ChatCompletionsRequest struct {
Model string `json:"model"`
*preprocessing.RenderJinjaTemplateRequest
*preprocessing.ApplyChatTemplateRequest
}

func main() {
Expand Down Expand Up @@ -320,35 +320,18 @@

logger.Info("Created ChatCompletions", "req", req)

// Get chat template for the model if not provided
if req.ChatTemplate == "" {
templateReq := preprocessing.FetchChatTemplateRequest{
Model: req.Model,
Token: os.Getenv(envHFToken),
}

var err error
req.ChatTemplate, req.ChatTemplateKWArgs, err = chatTemplatingProcessor.FetchChatTemplate(ctx, templateReq)
if err != nil {
http.Error(w, fmt.Sprintf("Failed to get chat template: %v", err), http.StatusInternalServerError)
return
}
}

response, err := chatTemplatingProcessor.RenderChatTemplate(ctx, req.RenderJinjaTemplateRequest)
renderedPrompt, err := chatTemplatingProcessor.ApplyChatTemplate(ctx, req.ApplyChatTemplateRequest)
if err != nil {
http.Error(w, fmt.Sprintf("Failed to render chat template: %v", err), http.StatusInternalServerError)
return
}

// Use KV-cache to score the rendered template
if len(response.RenderedChats) == 0 {
http.Error(w, "No rendered chats found in response", http.StatusInternalServerError)
if len(renderedPrompt) == 0 {

Check failure on line 330 in examples/kv_events/online/main.go

View workflow job for this annotation

GitHub Actions / lint-and-test

emptyStringTest: replace `len(renderedPrompt) == 0` with `renderedPrompt == ""` (gocritic)
http.Error(w, "rendered prompt is empty", http.StatusInternalServerError)
return
}

renderedPrompt := response.RenderedChats[0]

// Get score
pods, err := kvCacheIndexer.GetPodScores(ctx, nil, renderedPrompt, req.Model, nil)
if err != nil {
Expand Down
2 changes: 1 addition & 1 deletion examples/testdata/data.go
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ const (
ModelName = "bert-base-uncased"
)

var RenderReq *preprocessing.RenderJinjaTemplateRequest = nil
var RenderReq *preprocessing.ApplyChatTemplateRequest = nil

//go:embed prompt.txt
var Prompt string
Expand Down
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ require (
github.com/stretchr/testify v1.10.0
github.com/vmihailenco/msgpack/v5 v5.4.1
go.uber.org/multierr v1.11.0
go.uber.org/zap v1.27.0
Comment thread
hyeongyun0916 marked this conversation as resolved.
golang.org/x/net v0.38.0
google.golang.org/grpc v1.68.1
google.golang.org/protobuf v1.36.5
Expand Down Expand Up @@ -53,7 +54,6 @@ require (
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
github.com/x448/float16 v0.8.4 // indirect
github.com/yuin/gopher-lua v1.1.1 // indirect
go.uber.org/zap v1.27.0 // indirect
golang.org/x/oauth2 v0.27.0 // indirect
golang.org/x/sys v0.35.0 // indirect
golang.org/x/term v0.30.0 // indirect
Expand Down
2 changes: 1 addition & 1 deletion pkg/kvcache/indexer.go
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,7 @@
return nil, fmt.Errorf("failed to create KVBlockScorer: %w", err)
}

tokenizersPool, err := tokenization.NewTokenizationPool(config.TokenizersPoolConfig, tokensIndexer)

Check failure on line 105 in pkg/kvcache/indexer.go

View workflow job for this annotation

GitHub Actions / lint-and-test

Function `NewTokenizationPool->NewCachedLocalTokenizer` should pass the context parameter (contextcheck)
if err != nil {
return nil, fmt.Errorf("failed to create tokenizers pool: %w", err)
}
Expand Down Expand Up @@ -134,7 +134,7 @@
// relevant.
//
// The function returns a map of pod identifiers to scores.
func (k *Indexer) GetPodScores(ctx context.Context, renderReq *preprocessing.RenderJinjaTemplateRequest, prompt, modelName string,
func (k *Indexer) GetPodScores(ctx context.Context, renderReq *preprocessing.ApplyChatTemplateRequest, prompt, modelName string,
podIdentifiers []string,
) (map[string]float64, error) {
traceLogger := log.FromContext(ctx).V(logging.TRACE).WithName("kvcache.GetPodScores")
Expand Down
Loading
Loading