533 lines
14 KiB
Go
533 lines
14 KiB
Go
package autotune
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"math"
|
|
"net/http"
|
|
"net/url"
|
|
"os"
|
|
"sort"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/example/ollama-fair-gateway/internal/config"
|
|
"github.com/example/ollama-fair-gateway/internal/hoststats"
|
|
"github.com/example/ollama-fair-gateway/internal/state"
|
|
"github.com/example/ollama-fair-gateway/internal/worker"
|
|
)
|
|
|
|
type Sample struct {
|
|
TTFTMS float64 `json:"ttft_ms"`
|
|
ServiceMS float64 `json:"service_ms"`
|
|
PromptTokens int64 `json:"prompt_tokens"`
|
|
CompletionTokens int64 `json:"completion_tokens"`
|
|
PromptTPS float64 `json:"prompt_tps"`
|
|
OutputTPS float64 `json:"output_tps"`
|
|
Status int `json:"status"`
|
|
Error string `json:"error,omitempty"`
|
|
}
|
|
|
|
type Level struct {
|
|
Concurrency int `json:"concurrency"`
|
|
Requests int `json:"requests"`
|
|
Successful int `json:"successful"`
|
|
TTFTP50MS float64 `json:"ttft_p50_ms"`
|
|
TTFTP95MS float64 `json:"ttft_p95_ms"`
|
|
MeanServiceMS float64 `json:"mean_service_ms"`
|
|
MeanPromptTPS float64 `json:"mean_prompt_tps"`
|
|
MeanOutputTPS float64 `json:"mean_output_tps"`
|
|
AggregateOutputTPS float64 `json:"aggregate_output_tps"`
|
|
PeakVRAMBytes int64 `json:"peak_vram_bytes,omitempty"`
|
|
PeakGPUPercent float64 `json:"peak_gpu_percent,omitempty"`
|
|
Score float64 `json:"score"`
|
|
Samples []Sample `json:"samples,omitempty"`
|
|
}
|
|
|
|
type Profile struct {
|
|
ID string `json:"id"`
|
|
Worker string `json:"worker"`
|
|
Model string `json:"model"`
|
|
Status string `json:"status"` // queued|running|completed|failed|cancelled
|
|
StartedAt time.Time `json:"started_at"`
|
|
FinishedAt time.Time `json:"finished_at,omitempty"`
|
|
Levels []Level `json:"levels,omitempty"`
|
|
RecommendedConcurrency int `json:"recommended_concurrency,omitempty"`
|
|
Applied bool `json:"applied,omitempty"`
|
|
AppliedAt time.Time `json:"applied_at,omitempty"`
|
|
Error string `json:"error,omitempty"`
|
|
}
|
|
|
|
type persistentState struct {
|
|
Version int `json:"version"`
|
|
SavedAt time.Time `json:"saved_at"`
|
|
Profiles []Profile `json:"profiles"`
|
|
Applied map[string]map[string]int `json:"applied"`
|
|
}
|
|
|
|
type StartRequest struct {
|
|
Worker string `json:"worker"`
|
|
Model string `json:"model"`
|
|
MaxConcurrency int `json:"max_concurrency,omitempty"`
|
|
SamplesPerLevel int `json:"samples_per_level,omitempty"`
|
|
MaxTokens int `json:"max_tokens,omitempty"`
|
|
Prompt string `json:"prompt,omitempty"`
|
|
}
|
|
|
|
type Manager struct {
|
|
mu sync.RWMutex
|
|
cfg config.AutoTuningConfig
|
|
workers *worker.Pool
|
|
path string
|
|
client *http.Client
|
|
profiles map[string]Profile
|
|
order []string
|
|
applied map[string]map[string]int
|
|
cancel map[string]context.CancelFunc
|
|
}
|
|
|
|
func New(cfg config.AutoTuningConfig, workers *worker.Pool, path string) (*Manager, error) {
|
|
m := &Manager{cfg: cfg, workers: workers, path: path, client: &http.Client{Transport: &http.Transport{Proxy: http.ProxyFromEnvironment, MaxIdleConns: 64, MaxIdleConnsPerHost: 32, IdleConnTimeout: 30 * time.Second}}, profiles: map[string]Profile{}, applied: map[string]map[string]int{}, cancel: map[string]context.CancelFunc{}}
|
|
var ps persistentState
|
|
if path != "" {
|
|
err := (state.AtomicJSON{Path: path, Mode: 0640}).Load(&ps)
|
|
if err != nil && !errors.Is(err, os.ErrNotExist) {
|
|
return nil, err
|
|
}
|
|
}
|
|
for _, p := range ps.Profiles {
|
|
if p.ID == "" {
|
|
continue
|
|
}
|
|
if p.Status == "running" || p.Status == "queued" {
|
|
p.Status = "failed"
|
|
p.Error = "gateway restarted while benchmark was active"
|
|
p.FinishedAt = time.Now().UTC()
|
|
}
|
|
m.profiles[p.ID] = p
|
|
m.order = append(m.order, p.ID)
|
|
}
|
|
if ps.Applied != nil {
|
|
m.applied = ps.Applied
|
|
}
|
|
return m, nil
|
|
}
|
|
|
|
func (m *Manager) saveLocked() error {
|
|
if m.path == "" {
|
|
return nil
|
|
}
|
|
profiles := make([]Profile, 0, len(m.order))
|
|
for _, id := range m.order {
|
|
if p, ok := m.profiles[id]; ok {
|
|
profiles = append(profiles, p)
|
|
}
|
|
}
|
|
ps := persistentState{Version: 1, SavedAt: time.Now().UTC(), Profiles: profiles, Applied: m.applied}
|
|
return (state.AtomicJSON{Path: m.path, Mode: 0640}).Save(ps)
|
|
}
|
|
|
|
func (m *Manager) Applied() map[string]map[string]int {
|
|
m.mu.RLock()
|
|
defer m.mu.RUnlock()
|
|
out := map[string]map[string]int{}
|
|
for w, mm := range m.applied {
|
|
out[w] = map[string]int{}
|
|
for model, n := range mm {
|
|
out[w][model] = n
|
|
}
|
|
}
|
|
return out
|
|
}
|
|
|
|
func (m *Manager) List() []Profile {
|
|
m.mu.RLock()
|
|
defer m.mu.RUnlock()
|
|
out := make([]Profile, 0, len(m.order))
|
|
for i := len(m.order) - 1; i >= 0; i-- {
|
|
if p, ok := m.profiles[m.order[i]]; ok {
|
|
out = append(out, cloneProfile(p))
|
|
}
|
|
}
|
|
return out
|
|
}
|
|
func (m *Manager) Get(id string) (Profile, bool) {
|
|
m.mu.RLock()
|
|
defer m.mu.RUnlock()
|
|
p, ok := m.profiles[id]
|
|
return cloneProfile(p), ok
|
|
}
|
|
|
|
func cloneProfile(p Profile) Profile {
|
|
p.Levels = append([]Level(nil), p.Levels...)
|
|
for i := range p.Levels {
|
|
p.Levels[i].Samples = append([]Sample(nil), p.Levels[i].Samples...)
|
|
}
|
|
return p
|
|
}
|
|
|
|
func (m *Manager) Start(parent context.Context, in StartRequest) (Profile, error) {
|
|
if !m.cfg.Enabled {
|
|
return Profile{}, errors.New("auto tuning is disabled")
|
|
}
|
|
in.Worker = strings.TrimSpace(in.Worker)
|
|
in.Model = strings.TrimSpace(in.Model)
|
|
if in.Worker == "" || in.Model == "" {
|
|
return Profile{}, errors.New("worker and model are required")
|
|
}
|
|
wc, ok := m.workers.WorkerConfig(in.Worker)
|
|
if !ok {
|
|
return Profile{}, fmt.Errorf("unknown worker %q", in.Worker)
|
|
}
|
|
if in.MaxConcurrency <= 0 {
|
|
in.MaxConcurrency = m.cfg.MaxConcurrency
|
|
}
|
|
in.MaxConcurrency = min(in.MaxConcurrency, wc.MaxConcurrent)
|
|
if in.MaxConcurrency < 1 {
|
|
in.MaxConcurrency = 1
|
|
}
|
|
if in.SamplesPerLevel <= 0 {
|
|
in.SamplesPerLevel = m.cfg.SamplesPerLevel
|
|
}
|
|
if in.SamplesPerLevel > 20 {
|
|
in.SamplesPerLevel = 20
|
|
}
|
|
if in.MaxTokens <= 0 {
|
|
in.MaxTokens = m.cfg.MaxTokens
|
|
}
|
|
if in.MaxTokens > 4096 {
|
|
in.MaxTokens = 4096
|
|
}
|
|
if strings.TrimSpace(in.Prompt) == "" {
|
|
in.Prompt = m.cfg.Prompt
|
|
}
|
|
id := fmt.Sprintf("tune-%d", time.Now().UnixNano())
|
|
p := Profile{ID: id, Worker: in.Worker, Model: in.Model, Status: "queued", StartedAt: time.Now().UTC()}
|
|
m.mu.Lock()
|
|
m.profiles[id] = p
|
|
m.order = append(m.order, id)
|
|
_ = m.saveLocked()
|
|
m.mu.Unlock()
|
|
ctx, cancel := context.WithTimeout(context.WithoutCancel(parent), m.cfg.Timeout.Value())
|
|
m.mu.Lock()
|
|
m.cancel[id] = cancel
|
|
m.mu.Unlock()
|
|
go m.run(ctx, id, in, wc)
|
|
return p, nil
|
|
}
|
|
|
|
func (m *Manager) Cancel(id string) bool {
|
|
m.mu.RLock()
|
|
c := m.cancel[id]
|
|
m.mu.RUnlock()
|
|
if c == nil {
|
|
return false
|
|
}
|
|
c()
|
|
return true
|
|
}
|
|
|
|
func (m *Manager) Apply(id string) (Profile, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
p, ok := m.profiles[id]
|
|
if !ok {
|
|
return Profile{}, errors.New("benchmark profile not found")
|
|
}
|
|
if p.Status != "completed" || p.RecommendedConcurrency <= 0 {
|
|
return Profile{}, errors.New("benchmark has no applicable recommendation")
|
|
}
|
|
if m.applied[p.Worker] == nil {
|
|
m.applied[p.Worker] = map[string]int{}
|
|
}
|
|
m.applied[p.Worker][p.Model] = p.RecommendedConcurrency
|
|
p.Applied = true
|
|
p.AppliedAt = time.Now().UTC()
|
|
m.profiles[id] = p
|
|
if err := m.saveLocked(); err != nil {
|
|
return Profile{}, err
|
|
}
|
|
return cloneProfile(p), nil
|
|
}
|
|
func (m *Manager) Reset(workerName, model string) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
if mm := m.applied[workerName]; mm != nil {
|
|
delete(mm, model)
|
|
if len(mm) == 0 {
|
|
delete(m.applied, workerName)
|
|
}
|
|
}
|
|
return m.saveLocked()
|
|
}
|
|
|
|
func (m *Manager) run(ctx context.Context, id string, in StartRequest, wc config.WorkerConfig) {
|
|
m.update(id, func(p *Profile) { p.Status = "running" })
|
|
levels := make([]Level, 0, in.MaxConcurrency)
|
|
for c := 1; c <= in.MaxConcurrency; c++ {
|
|
if ctx.Err() != nil {
|
|
m.finishCancelled(id, ctx.Err())
|
|
return
|
|
}
|
|
lvl := m.runLevel(ctx, wc, in.Model, in.Prompt, in.MaxTokens, c, in.SamplesPerLevel)
|
|
levels = append(levels, lvl)
|
|
m.update(id, func(p *Profile) { p.Levels = append([]Level(nil), levels...) })
|
|
}
|
|
if len(levels) == 0 {
|
|
m.finishFailed(id, "no benchmark levels completed")
|
|
return
|
|
}
|
|
maxAgg, maxP95 := 0.0, 0.0
|
|
for _, l := range levels {
|
|
if l.AggregateOutputTPS > maxAgg {
|
|
maxAgg = l.AggregateOutputTPS
|
|
}
|
|
if l.TTFTP95MS > maxP95 {
|
|
maxP95 = l.TTFTP95MS
|
|
}
|
|
}
|
|
best, bestScore := 1, -math.MaxFloat64
|
|
for i := range levels {
|
|
throughput := 0.0
|
|
if maxAgg > 0 {
|
|
throughput = levels[i].AggregateOutputTPS / maxAgg
|
|
}
|
|
latency := 0.0
|
|
if maxP95 > 0 {
|
|
latency = levels[i].TTFTP95MS / maxP95
|
|
}
|
|
failure := 1.0 - float64(levels[i].Successful)/float64(max(1, levels[i].Requests))
|
|
levels[i].Score = m.cfg.ThroughputWeight*throughput - m.cfg.TTFTWeight*latency - failure
|
|
if levels[i].Successful == levels[i].Requests && levels[i].Score > bestScore {
|
|
bestScore = levels[i].Score
|
|
best = levels[i].Concurrency
|
|
}
|
|
}
|
|
m.mu.Lock()
|
|
p := m.profiles[id]
|
|
p.Status = "completed"
|
|
p.Levels = levels
|
|
p.RecommendedConcurrency = best
|
|
p.FinishedAt = time.Now().UTC()
|
|
m.profiles[id] = p
|
|
delete(m.cancel, id)
|
|
_ = m.saveLocked()
|
|
m.mu.Unlock()
|
|
}
|
|
|
|
func (m *Manager) update(id string, fn func(*Profile)) {
|
|
m.mu.Lock()
|
|
p := m.profiles[id]
|
|
fn(&p)
|
|
m.profiles[id] = p
|
|
_ = m.saveLocked()
|
|
m.mu.Unlock()
|
|
}
|
|
func (m *Manager) finishCancelled(id string, err error) {
|
|
m.mu.Lock()
|
|
p := m.profiles[id]
|
|
p.Status = "cancelled"
|
|
if err != nil {
|
|
p.Error = err.Error()
|
|
}
|
|
p.FinishedAt = time.Now().UTC()
|
|
m.profiles[id] = p
|
|
delete(m.cancel, id)
|
|
_ = m.saveLocked()
|
|
m.mu.Unlock()
|
|
}
|
|
func (m *Manager) finishFailed(id, msg string) {
|
|
m.mu.Lock()
|
|
p := m.profiles[id]
|
|
p.Status = "failed"
|
|
p.Error = msg
|
|
p.FinishedAt = time.Now().UTC()
|
|
m.profiles[id] = p
|
|
delete(m.cancel, id)
|
|
_ = m.saveLocked()
|
|
m.mu.Unlock()
|
|
}
|
|
|
|
func (m *Manager) runLevel(ctx context.Context, wc config.WorkerConfig, model, prompt string, maxTokens, concurrency, repeats int) Level {
|
|
lvl := Level{Concurrency: concurrency, Requests: concurrency * repeats}
|
|
var samples []Sample
|
|
var totalTokens int64
|
|
var totalWall time.Duration
|
|
var peakVRAM int64
|
|
var peakGPU float64
|
|
var peakMu sync.Mutex
|
|
for rep := 0; rep < repeats; rep++ {
|
|
start := time.Now()
|
|
ch := make(chan Sample, concurrency)
|
|
var wg sync.WaitGroup
|
|
for i := 0; i < concurrency; i++ {
|
|
wg.Add(1)
|
|
go func() { defer wg.Done(); ch <- m.one(ctx, wc.URL, model, prompt, maxTokens) }()
|
|
}
|
|
done := make(chan struct{})
|
|
go func() { wg.Wait(); close(ch); close(done) }()
|
|
var telemetryWG sync.WaitGroup
|
|
if wc.NVIDIASMI {
|
|
telemetryWG.Add(1)
|
|
go func() {
|
|
defer telemetryWG.Done()
|
|
t := time.NewTicker(200 * time.Millisecond)
|
|
defer t.Stop()
|
|
for {
|
|
select {
|
|
case <-done:
|
|
return
|
|
case <-ctx.Done():
|
|
return
|
|
case <-t.C:
|
|
c, cancel := context.WithTimeout(ctx, time.Second)
|
|
n, e := hoststats.ReadNVIDIA(c, wc.NVIDIAGPU)
|
|
cancel()
|
|
if e == nil {
|
|
peakMu.Lock()
|
|
if n.MemoryUsedBytes > peakVRAM {
|
|
peakVRAM = n.MemoryUsedBytes
|
|
}
|
|
if n.UtilizationPercent > peakGPU {
|
|
peakGPU = n.UtilizationPercent
|
|
}
|
|
peakMu.Unlock()
|
|
}
|
|
}
|
|
}
|
|
}()
|
|
}
|
|
for s := range ch {
|
|
samples = append(samples, s)
|
|
if s.Status >= 200 && s.Status < 300 && s.Error == "" {
|
|
lvl.Successful++
|
|
totalTokens += s.CompletionTokens
|
|
}
|
|
}
|
|
telemetryWG.Wait()
|
|
totalWall += time.Since(start)
|
|
}
|
|
lvl.Samples = samples
|
|
lvl.PeakVRAMBytes = peakVRAM
|
|
lvl.PeakGPUPercent = peakGPU
|
|
var ttfts, services []float64
|
|
var pTPS, oTPS float64
|
|
for _, s := range samples {
|
|
if s.Error != "" || s.Status < 200 || s.Status >= 300 {
|
|
continue
|
|
}
|
|
ttfts = append(ttfts, s.TTFTMS)
|
|
services = append(services, s.ServiceMS)
|
|
pTPS += s.PromptTPS
|
|
oTPS += s.OutputTPS
|
|
}
|
|
lvl.TTFTP50MS = percentile(ttfts, .50)
|
|
lvl.TTFTP95MS = percentile(ttfts, .95)
|
|
lvl.MeanServiceMS = mean(services)
|
|
if lvl.Successful > 0 {
|
|
lvl.MeanPromptTPS = pTPS / float64(lvl.Successful)
|
|
lvl.MeanOutputTPS = oTPS / float64(lvl.Successful)
|
|
}
|
|
if totalWall > 0 {
|
|
lvl.AggregateOutputTPS = float64(totalTokens) / totalWall.Seconds()
|
|
}
|
|
return lvl
|
|
}
|
|
|
|
func (m *Manager) one(ctx context.Context, base, model, prompt string, maxTokens int) Sample {
|
|
u, err := url.Parse(strings.TrimRight(base, "/") + "/api/generate")
|
|
if err != nil {
|
|
return Sample{Error: err.Error()}
|
|
}
|
|
b, _ := json.Marshal(map[string]any{"model": model, "prompt": prompt, "stream": true, "keep_alive": "5m", "options": map[string]any{"num_predict": maxTokens, "temperature": 0}})
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, u.String(), bytes.NewReader(b))
|
|
if err != nil {
|
|
return Sample{Error: err.Error()}
|
|
}
|
|
req.Header.Set("Content-Type", "application/json")
|
|
start := time.Now()
|
|
resp, err := m.client.Do(req)
|
|
if err != nil {
|
|
return Sample{Error: err.Error(), ServiceMS: float64(time.Since(start).Microseconds()) / 1000}
|
|
}
|
|
defer resp.Body.Close()
|
|
out := Sample{Status: resp.StatusCode}
|
|
scan := bufio.NewScanner(resp.Body)
|
|
buf := make([]byte, 64*1024)
|
|
scan.Buffer(buf, 4<<20)
|
|
first := true
|
|
for scan.Scan() {
|
|
if first {
|
|
out.TTFTMS = float64(time.Since(start).Microseconds()) / 1000
|
|
first = false
|
|
}
|
|
var v map[string]any
|
|
if json.Unmarshal(scan.Bytes(), &v) != nil {
|
|
continue
|
|
}
|
|
if n := asI64(v["prompt_eval_count"]); n > 0 {
|
|
out.PromptTokens = n
|
|
}
|
|
if n := asI64(v["eval_count"]); n > 0 {
|
|
out.CompletionTokens = n
|
|
}
|
|
pn := asI64(v["prompt_eval_duration"])
|
|
en := asI64(v["eval_duration"])
|
|
if pn > 0 && out.PromptTokens > 0 {
|
|
out.PromptTPS = float64(out.PromptTokens) / (float64(pn) / 1e9)
|
|
}
|
|
if en > 0 && out.CompletionTokens > 0 {
|
|
out.OutputTPS = float64(out.CompletionTokens) / (float64(en) / 1e9)
|
|
}
|
|
}
|
|
if err := scan.Err(); err != nil {
|
|
out.Error = err.Error()
|
|
}
|
|
out.ServiceMS = float64(time.Since(start).Microseconds()) / 1000
|
|
return out
|
|
}
|
|
func asI64(v any) int64 {
|
|
switch x := v.(type) {
|
|
case float64:
|
|
return int64(x)
|
|
case json.Number:
|
|
n, _ := x.Int64()
|
|
return n
|
|
case int64:
|
|
return x
|
|
case int:
|
|
return int64(x)
|
|
}
|
|
return 0
|
|
}
|
|
func mean(xs []float64) float64 {
|
|
if len(xs) == 0 {
|
|
return 0
|
|
}
|
|
s := 0.0
|
|
for _, x := range xs {
|
|
s += x
|
|
}
|
|
return s / float64(len(xs))
|
|
}
|
|
func percentile(xs []float64, p float64) float64 {
|
|
if len(xs) == 0 {
|
|
return 0
|
|
}
|
|
ys := append([]float64(nil), xs...)
|
|
sort.Float64s(ys)
|
|
i := int(math.Ceil(p*float64(len(ys)))) - 1
|
|
if i < 0 {
|
|
i = 0
|
|
}
|
|
if i >= len(ys) {
|
|
i = len(ys) - 1
|
|
}
|
|
return ys[i]
|
|
}
|