llm: preserve generation headroom for shifted prompts (#16856)
--------- Co-authored-by: Daniel Hiltgen <daniel@ollama.com>
This commit is contained in:
@@ -274,8 +274,8 @@ func (s *llamaServerRunner) completionPromptForRequest(ctx context.Context, req
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
limit := s.options.NumCtx - 1
|
fullPromptLimit := s.options.NumCtx - 1
|
||||||
if len(tokens) <= limit {
|
if len(tokens) <= fullPromptLimit {
|
||||||
return prompt, nil
|
return prompt, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -290,8 +290,12 @@ func (s *llamaServerRunner) completionPromptForRequest(ctx context.Context, req
|
|||||||
if nKeep < 0 {
|
if nKeep < 0 {
|
||||||
nKeep = len(tokens)
|
nKeep = len(tokens)
|
||||||
}
|
}
|
||||||
nKeep = min(nKeep, limit)
|
if s.tokenizerAddsBOS() {
|
||||||
|
nKeep++
|
||||||
|
}
|
||||||
|
nKeep = min(nKeep, fullPromptLimit)
|
||||||
|
|
||||||
|
limit := contextShiftPromptLimit(s.options.NumCtx, nKeep)
|
||||||
discard := len(tokens) - limit
|
discard := len(tokens) - limit
|
||||||
truncated := make([]int, 0, limit)
|
truncated := make([]int, 0, limit)
|
||||||
truncated = append(truncated, tokens[:nKeep]...)
|
truncated = append(truncated, tokens[:nKeep]...)
|
||||||
@@ -301,6 +305,18 @@ func (s *llamaServerRunner) completionPromptForRequest(ctx context.Context, req
|
|||||||
return truncated, nil
|
return truncated, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func contextShiftPromptLimit(numCtx, numKeep int) int {
|
||||||
|
if numCtx <= 1 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
numKeep = max(0, min(numKeep, numCtx-1))
|
||||||
|
|
||||||
|
// Match the old runners' first context shift: preserve num_keep, then free
|
||||||
|
// roughly half of the remaining context before generation needs the slot.
|
||||||
|
return numCtx - max((numCtx-numKeep)/2, 1)
|
||||||
|
}
|
||||||
|
|
||||||
func (s *llamaServerRunner) ContextLength() int {
|
func (s *llamaServerRunner) ContextLength() int {
|
||||||
return s.options.NumCtx
|
return s.options.NumCtx
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -127,6 +127,30 @@ func TestBoundedNumPredict(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestContextShiftPromptLimit(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
numCtx int
|
||||||
|
numKeep int
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{name: "small context reserves half after keep", numCtx: 8, numKeep: 3, want: 6},
|
||||||
|
{name: "issue 16618 context preserves generation headroom", numCtx: 4096, numKeep: 4, want: 2050},
|
||||||
|
{name: "issue 16618 with implicit BOS keep", numCtx: 4096, numKeep: 5, want: 2051},
|
||||||
|
{name: "keep is clamped below context", numCtx: 8, numKeep: 99, want: 7},
|
||||||
|
{name: "negative keep is treated as zero", numCtx: 8, numKeep: -1, want: 4},
|
||||||
|
{name: "invalid context has no prompt budget", numCtx: 1, numKeep: 0, want: 0},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := contextShiftPromptLimit(tt.numCtx, tt.numKeep); got != tt.want {
|
||||||
|
t.Fatalf("contextShiftPromptLimit(%d, %d) = %d, want %d", tt.numCtx, tt.numKeep, got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLlamaServerCompletionSSEParsing(t *testing.T) {
|
func TestLlamaServerCompletionSSEParsing(t *testing.T) {
|
||||||
// Simulate llama-server SSE streaming response
|
// Simulate llama-server SSE streaming response
|
||||||
sseLines := []string{
|
sseLines := []string{
|
||||||
@@ -654,7 +678,7 @@ func TestLlamaServerCompletionContextShiftTruncatesPromptOverContext(t *testing.
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
w.Header().Set("Content-Type", "text/event-stream")
|
w.Header().Set("Content-Type", "text/event-stream")
|
||||||
fmt.Fprintln(w, `data: {"content":"ok","stop":true,"timings":{"prompt_n":7,"prompt_ms":1,"predicted_n":1,"predicted_ms":1}}`)
|
fmt.Fprintln(w, `data: {"content":"ok","stop":true,"timings":{"prompt_n":6,"prompt_ms":1,"predicted_n":1,"predicted_ms":1}}`)
|
||||||
default:
|
default:
|
||||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||||
}
|
}
|
||||||
@@ -698,7 +722,7 @@ func TestLlamaServerCompletionContextShiftTruncatesPromptOverContext(t *testing.
|
|||||||
if !ok {
|
if !ok {
|
||||||
t.Fatalf("completion prompt = %T, want token array", capturedReq.Prompt)
|
t.Fatalf("completion prompt = %T, want token array", capturedReq.Prompt)
|
||||||
}
|
}
|
||||||
want := []int{0, 1, 2, 6, 7, 8, 9}
|
want := []int{0, 1, 2, 7, 8, 9}
|
||||||
if len(got) != len(want) {
|
if len(got) != len(want) {
|
||||||
t.Fatalf("token prompt len = %d, want %d: %#v", len(got), len(want), got)
|
t.Fatalf("token prompt len = %d, want %d: %#v", len(got), len(want), got)
|
||||||
}
|
}
|
||||||
@@ -713,6 +737,99 @@ func TestLlamaServerCompletionContextShiftTruncatesPromptOverContext(t *testing.
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLlamaServerCompletionContextShiftAvoidsOneTokenHeadroomRegression(t *testing.T) {
|
||||||
|
var capturedReq llamaServerCompletionRequest
|
||||||
|
var tokenizeReq struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
AddSpecial bool `json:"add_special"`
|
||||||
|
ParseSpecial *bool `json:"parse_special"`
|
||||||
|
}
|
||||||
|
tokens := make([]int, 5000)
|
||||||
|
for i := range tokens {
|
||||||
|
tokens[i] = i
|
||||||
|
}
|
||||||
|
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
switch r.URL.Path {
|
||||||
|
case "/health":
|
||||||
|
fmt.Fprint(w, `{"status":"ok"}`)
|
||||||
|
case "/tokenize":
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&tokenizeReq); err != nil {
|
||||||
|
t.Errorf("invalid tokenize request body: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := json.NewEncoder(w).Encode(map[string][]int{"tokens": tokens}); err != nil {
|
||||||
|
t.Errorf("failed to encode tokenize response: %v", err)
|
||||||
|
}
|
||||||
|
case "/completion":
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&capturedReq); err != nil {
|
||||||
|
t.Errorf("invalid completion request body: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "text/event-stream")
|
||||||
|
fmt.Fprintln(w, `data: {"content":"ok","stop":true,"timings":{"prompt_n":2051,"prompt_ms":1,"predicted_n":32,"predicted_ms":1}}`)
|
||||||
|
default:
|
||||||
|
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
parts := strings.Split(srv.URL, ":")
|
||||||
|
var portInt int
|
||||||
|
fmt.Sscanf(parts[len(parts)-1], "%d", &portInt)
|
||||||
|
|
||||||
|
runner := &llamaServerRunner{
|
||||||
|
port: portInt,
|
||||||
|
cmd: fakeRunningCmd(),
|
||||||
|
sem: semaphore.NewWeighted(1),
|
||||||
|
options: api.Options{Runner: api.Runner{NumCtx: 4096}},
|
||||||
|
ggml: loadTestGGML(t, ggml.KV{
|
||||||
|
"general.architecture": "gemma3",
|
||||||
|
"tokenizer.ggml.add_bos_token": true,
|
||||||
|
}),
|
||||||
|
launch: llamaServerLaunchConfig{
|
||||||
|
config: LlamaServerConfig{ContextShift: true},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := api.DefaultOptions()
|
||||||
|
opts.NumKeep = 4
|
||||||
|
err := runner.Completion(t.Context(), CompletionRequest{
|
||||||
|
Prompt: strings.Repeat("long prompt ", 500),
|
||||||
|
Options: &opts,
|
||||||
|
Truncate: true,
|
||||||
|
}, func(cr CompletionResponse) {})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Completion error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, ok := capturedReq.Prompt.([]any)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("completion prompt = %T, want token array", capturedReq.Prompt)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(got) != 2051 {
|
||||||
|
t.Fatalf("token prompt len = %d, want 2051", len(got))
|
||||||
|
}
|
||||||
|
if len(got) == 4095 {
|
||||||
|
t.Fatal("token prompt preserved old one-token headroom behavior")
|
||||||
|
}
|
||||||
|
|
||||||
|
effectiveKeep := opts.NumKeep + 1
|
||||||
|
for i := range effectiveKeep {
|
||||||
|
gotToken, ok := got[i].(float64)
|
||||||
|
if !ok || int(gotToken) != i {
|
||||||
|
t.Fatalf("token prompt[%d] = %#v, want %d", i, got[i], i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const wantSuffixStart = 2954
|
||||||
|
gotToken, ok := got[effectiveKeep].(float64)
|
||||||
|
if !ok || int(gotToken) != wantSuffixStart {
|
||||||
|
t.Fatalf("first shifted token = %#v, want %d", got[effectiveKeep], wantSuffixStart)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLlamaServerCompletionWithMediaUsesRunnerMarker(t *testing.T) {
|
func TestLlamaServerCompletionWithMediaUsesRunnerMarker(t *testing.T) {
|
||||||
var capturedReq llamaServerCompletionRequest
|
var capturedReq llamaServerCompletionRequest
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user