mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Refactor ollama to distinguish loading models and prompting as separate.
This takes advantage of new capabilities added in recent versions of ollama to do model-loading-only prompt requests. PiperOrigin-RevId: 612149945
This commit is contained in:
committed by
gVisor bot
parent
f7dd4ceebe
commit
1ed73f1f6f
@@ -6,7 +6,7 @@ ENV OLLAMA_ORIGINS=*
|
||||
ENV OLLAMA_HOST=0.0.0.0:11434
|
||||
|
||||
# Pre-install models useful for benchmarking.
|
||||
# These are huge (total ~100 GiB), but necessary to benchmark
|
||||
# These are huge (total ~115 GiB), but necessary to benchmark
|
||||
# models of various sizes. They are in their own image file to
|
||||
# keep the test-only image lighter by comparison.
|
||||
RUN bash -c ' \
|
||||
@@ -21,6 +21,7 @@ RUN bash -c ' \
|
||||
ollama pull mixtral:instruct && \
|
||||
ollama pull gemma:2b-instruct && \
|
||||
ollama pull gemma:7b-instruct && \
|
||||
ollama pull llava:7b-v1.6 && \
|
||||
kill "$serverpid" && \
|
||||
wait "$serverpid" \
|
||||
'
|
||||
|
||||
+60
-69
@@ -107,11 +107,7 @@ func New(ctx context.Context, server Server, logger testutil.Logger) (*Ollama, e
|
||||
// we cannot detect if it is using the GPU or not.
|
||||
// This may fail during the process of loading the first model, so we keep
|
||||
// iterating for a while.
|
||||
_, err = llm.Prompt(ctx, &Prompt{
|
||||
Model: &Model{Name: llm.ModelNames[0]},
|
||||
Preamble: DoNothing,
|
||||
Query: curtQuery,
|
||||
})
|
||||
_, err = llm.WarmModel(ctx, &Model{Name: llm.ModelNames[0]}, 1*time.Millisecond, false)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("could not load first model %q: %w", llm.ModelNames[0], err)
|
||||
}
|
||||
@@ -425,27 +421,17 @@ func ZeroTemperatureModel(name string) *Model {
|
||||
}
|
||||
}
|
||||
|
||||
// Preamble dictates what to do before actually issuing a prompt.
|
||||
type Preamble int
|
||||
|
||||
const (
|
||||
// DoNothing means to do nothing when issuing a prompt.
|
||||
DoNothing Preamble = iota
|
||||
|
||||
// WarmUp issues a short query to the same model before
|
||||
// issuing a prompt. This ensures the model is warm.
|
||||
WarmUp Preamble = iota
|
||||
|
||||
// Unload issues a short query with a different model before
|
||||
// issuing a prompt. This ensures the model is cold.
|
||||
Unload Preamble = iota
|
||||
)
|
||||
|
||||
// Prompt is an ollama prompt.
|
||||
type Prompt struct {
|
||||
// Model is the model to query.
|
||||
Model *Model
|
||||
|
||||
// If set, keep the model alive in memory for the given duration after this
|
||||
// prompt is answered. A zero duration will use the ollama default (a few
|
||||
// minutes). Note that model unloading is asynchronous, so the model will
|
||||
// not be fully unloaded after only `KeepModelAlive` beyond prompt response.
|
||||
KeepModelAlive time.Duration
|
||||
|
||||
// Query is the prompt string.
|
||||
// Common leading whitespace will be removed.
|
||||
Query string
|
||||
@@ -453,11 +439,6 @@ type Prompt struct {
|
||||
// Context is the conversational context to follow up on, if any.
|
||||
// This is returned from `Response`.
|
||||
Context ConversationContext
|
||||
|
||||
// Preamble dictates what to do before actually issuing this prompt.
|
||||
// This can be used to make sure the model is warm, or to make sure
|
||||
// the model is cold.
|
||||
Preamble Preamble
|
||||
}
|
||||
|
||||
// CleanQuery removes common whitespace from query lines, and all
|
||||
@@ -546,21 +527,27 @@ func (p *Prompt) WithHotterModel() *Prompt {
|
||||
|
||||
// PromptJSON encodes the JSON data for a query.
|
||||
type PromptJSON struct {
|
||||
Model string `json:"model"`
|
||||
Prompt string `json:"prompt"`
|
||||
Stream bool `json:"stream"`
|
||||
Context ConversationContext `json:"context"`
|
||||
Options map[string]any `json:"options"`
|
||||
Model string `json:"model"`
|
||||
Prompt string `json:"prompt,omitempty"`
|
||||
Stream bool `json:"stream"`
|
||||
Context ConversationContext `json:"context"`
|
||||
Options map[string]any `json:"options"`
|
||||
KeepAlive string `json:"keep_alive,omitempty"`
|
||||
}
|
||||
|
||||
// json encodes this prompt to the JSON format expected by Ollama.
|
||||
func (p *Prompt) json() PromptJSON {
|
||||
keepAlive := ""
|
||||
if p.KeepModelAlive != 0 {
|
||||
keepAlive = p.KeepModelAlive.String()
|
||||
}
|
||||
return PromptJSON{
|
||||
Model: p.Model.Name,
|
||||
Prompt: p.CleanQuery(),
|
||||
Stream: true,
|
||||
Context: p.Context,
|
||||
Options: p.Model.Options,
|
||||
Model: p.Model.Name,
|
||||
Prompt: p.CleanQuery(),
|
||||
Stream: true,
|
||||
Context: p.Context,
|
||||
Options: p.Model.Options,
|
||||
KeepAlive: keepAlive,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -732,7 +719,8 @@ func (r *Response) TotalDuration() time.Duration {
|
||||
return time.Duration(r.data[len(r.data)-1].TotalNanos) * time.Nanosecond
|
||||
}
|
||||
|
||||
// LoadDuration returns the load response generation time.
|
||||
// LoadDuration returns the load response generation time as reported
|
||||
// by the ollama server.
|
||||
func (r *Response) LoadDuration() time.Duration {
|
||||
if !r.Done() {
|
||||
return time.Duration(0)
|
||||
@@ -794,35 +782,47 @@ func (llm *Ollama) getReplacementModel(model *Model) (*Model, error) {
|
||||
return nil, fmt.Errorf("cannot find a replacement model to load instead of %q (available: %v; cheap: %v)", model.Name, llm.ModelNames, llm.cheapModels)
|
||||
}
|
||||
|
||||
// Prompt returns the result of prompting the given `model` with `prompt`.
|
||||
func (llm *Ollama) Prompt(ctx context.Context, prompt *Prompt) (*Response, error) {
|
||||
switch prompt.Preamble {
|
||||
case DoNothing: // Do nothing.
|
||||
case WarmUp:
|
||||
warmCtx, warmCancel := context.WithTimeout(ctx, 3*time.Minute)
|
||||
_, err := jsonPost[PromptJSON, ResponseJSON](warmCtx, llm, "/api/generate", (&Prompt{
|
||||
Model: prompt.Model,
|
||||
Query: curtQuery,
|
||||
}).json())
|
||||
warmCancel()
|
||||
// ModelLoadStats holds metrics about the model loading process.
|
||||
type ModelLoadStats struct {
|
||||
// ClientReportedDuration is the duration to load the model as perceived
|
||||
// by the client, measured by HTTP client metrics.
|
||||
ClientReportedDuration time.Duration
|
||||
}
|
||||
|
||||
// WarmModel pre-warms a model in memory and keeps it warm for `keepWarmFor`.
|
||||
// If `unloadFirst` is true, another model will be loaded before loading the
|
||||
// requested model. This ensures that the model was loaded from a cold state.
|
||||
func (llm *Ollama) WarmModel(ctx context.Context, model *Model, keepWarmFor time.Duration, unloadFirst bool) (*ModelLoadStats, error) {
|
||||
if keepWarmFor <= 0 {
|
||||
return nil, fmt.Errorf("keepWarmFor must be strictly positive, got %v", keepWarmFor)
|
||||
}
|
||||
if unloadFirst {
|
||||
replacementModel, err := llm.getReplacementModel(model)
|
||||
if err != nil {
|
||||
return nil, llm.withServerLogsErr(ctx, fmt.Errorf("warmup prompt for model %s failed: %w", prompt.Model.Name, err))
|
||||
}
|
||||
case Unload:
|
||||
replacementModel, err := llm.getReplacementModel(prompt.Model)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cannot find a replacement model to load instead of %q to forcefully unload it: %w", prompt.Model.Name, err)
|
||||
return nil, fmt.Errorf("cannot find a replacement model to load instead of %q to forcefully unload it: %w", model.Name, err)
|
||||
}
|
||||
unloadCtx, unloadCancel := context.WithTimeout(ctx, 3*time.Minute)
|
||||
_, err = jsonPost[PromptJSON, ResponseJSON](unloadCtx, llm, "/api/generate", (&Prompt{
|
||||
Model: replacementModel,
|
||||
Query: curtQuery,
|
||||
}).json())
|
||||
_, err = llm.Prompt(unloadCtx, &Prompt{Model: replacementModel, KeepModelAlive: 1 * time.Millisecond})
|
||||
unloadCancel()
|
||||
if err != nil {
|
||||
return nil, llm.withServerLogsErr(ctx, fmt.Errorf("unload prompt for model %s failed: %w", replacementModel.Name, err))
|
||||
return nil, llm.withServerLogsErr(ctx, fmt.Errorf("unload prompt for replacement model %s failed: %w", replacementModel.Name, err))
|
||||
}
|
||||
select { // Wait for the model to get unloaded. Unfortunately there isn't a great way to do this but to sleep.
|
||||
case <-time.After(20 * time.Second):
|
||||
case <-ctx.Done():
|
||||
}
|
||||
}
|
||||
resp, err := llm.Prompt(ctx, &Prompt{Model: model, KeepModelAlive: keepWarmFor})
|
||||
if err != nil {
|
||||
return nil, llm.withServerLogsErr(ctx, fmt.Errorf("warmup prompt for model %s failed: %w", model.Name, err))
|
||||
}
|
||||
return &ModelLoadStats{
|
||||
ClientReportedDuration: resp.metrics.LastByteRead.Sub(resp.metrics.RequestSent),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Prompt returns the result of prompting the given `model` with `prompt`.
|
||||
func (llm *Ollama) Prompt(ctx context.Context, prompt *Prompt) (*Response, error) {
|
||||
resp, err := jsonPost[PromptJSON, ResponseJSON](ctx, llm, "/api/generate", prompt.json())
|
||||
if err != nil {
|
||||
return nil, llm.withServerLogsErr(ctx, fmt.Errorf("prompt (%s %q) request failed: %w", prompt.Model.Name, prompt.CleanQuery(), err))
|
||||
@@ -832,25 +832,16 @@ func (llm *Ollama) Prompt(ctx context.Context, prompt *Prompt) (*Response, error
|
||||
|
||||
// PromptUntil repeatedly issues a prompt until `iterate` returns a nil error.
|
||||
// `iterate` may optionally return an updated `Prompt` which will be used to
|
||||
// follow up.
|
||||
// This is useful to work around the flakiness of LLMs in tests.
|
||||
// follow up. This is useful to work around the flakiness of LLMs in tests.
|
||||
func (llm *Ollama) PromptUntil(ctx context.Context, prompt *Prompt, iterate func(*Prompt, *Response) (*Prompt, error)) (*Response, error) {
|
||||
var lastResponse *Response
|
||||
var lastError error
|
||||
attempts := 0
|
||||
warmed := false
|
||||
for ctx.Err() == nil {
|
||||
response, err := llm.Prompt(ctx, prompt)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prompt request failed: %w", err)
|
||||
}
|
||||
if prompt.Preamble == WarmUp && !warmed {
|
||||
// Future prompts do not need to specify the WarmFirst option.
|
||||
promptCopy := *prompt
|
||||
promptCopy.Preamble = DoNothing
|
||||
prompt = &promptCopy
|
||||
warmed = true
|
||||
}
|
||||
attempts++
|
||||
newPrompt, err := iterate(prompt, response)
|
||||
if err == nil {
|
||||
|
||||
@@ -89,7 +89,6 @@ func TestLLM(t *testing.T) {
|
||||
Translate the following text from English to Chinese:
|
||||
"Hello World".
|
||||
`,
|
||||
Preamble: ollama.DoNothing,
|
||||
}
|
||||
promptCtx, promptCancel := context.WithTimeout(ctx, 3*time.Minute)
|
||||
response, err := llm.PromptUntil(promptCtx, &prompt, func(prompt *ollama.Prompt, response *ollama.Response) (*ollama.Prompt, error) {
|
||||
@@ -147,7 +146,6 @@ func TestLLM(t *testing.T) {
|
||||
the code works and do not provide usage examples.
|
||||
Output a single block of Python code wrapped between %q marks.
|
||||
`, markerString, codeBlockDelim),
|
||||
Preamble: ollama.DoNothing,
|
||||
}
|
||||
response, err := llm.PromptUntil(promptCtx, &prompt, func(prompt *ollama.Prompt, response *ollama.Response) (*ollama.Prompt, error) {
|
||||
defer prompt.Model.RaiseTemperature()
|
||||
|
||||
Reference in New Issue
Block a user