From a52a3cd0e6d6d999e93877150adf4c4ec19deff0 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 25 Mar 2026 21:50:15 +0100 Subject: Refactor llm module --- llm/perplexity.go | 159 ------------------------------------------------- llm/perplexity_test.go | 45 -------------- llm/run.go | 159 +++++++++++++++++++++++++++++++++++++++++++++++++ llm/run_test.go | 45 ++++++++++++++ 4 files changed, 204 insertions(+), 204 deletions(-) delete mode 100644 llm/perplexity.go delete mode 100644 llm/perplexity_test.go create mode 100644 llm/run.go create mode 100644 llm/run_test.go (limited to 'llm') diff --git a/llm/perplexity.go b/llm/perplexity.go deleted file mode 100644 index df48ff6..0000000 --- a/llm/perplexity.go +++ /dev/null @@ -1,159 +0,0 @@ -package llm - -import ( - "fmt" - "sync" - "time" - - "go.jknobloc.com/x/dataset" - "go.jknobloc.com/x/tui" -) - -func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int) error { - devices := make([]int, len(e.models)) - - for i := range len(devices) { - devices[i] = i - } - - devicePool := newPool(devices...) - - var wg sync.WaitGroup - - for range e.numWorkers { - wg.Add(1) - - go func() { - defer wg.Done() - - for b := range e.jobs { - device := devicePool.Acquire() - - func() { - defer devicePool.Release(device) - - defer func() { - if r := recover(); r != nil { - fmt.Println("HOUSTON") // TODO handle - } - }() - - e.execute(&b, device) - }() - - e.completed.Add(int64(b.Size())) - } - }() - } - - pb := tui.NewProgressBar(title, 20, 0, time.Now()) - - pb.Start(1*time.Second, func() int { - return int(e.completed.Load()) - }) - - defer pb.Close() - - n := 0 - - for d := range data.Texts("text") { - tokens := toInt64(e.tokenizer.Tokenize(d)) - - e.schedule(n, tokens, window, stride, e.batchSize) - - pb.SetTotal(pb.Total() + e.estimateJobs(tokens, window, stride)) - - n++ - } - - close(e.jobs) - - wg.Wait() - - close(e.results) - - return nil -} - -func (e *Evaluator[R]) estimateJobs(tokens []int64, window, stride int) int { - if len(tokens) < window { - return 0 - } - - windows := ((len(tokens) - window) / stride) + 1 - // jobs := (windows + batchSize - 1) / e.batchSize - - return windows -} - -func (e *Evaluator[R]) schedule(uid int, tokens []int64, contextSize, stride, batchSize int) { - b := newBatch(batchSize) - - seen := 1 // first token as context - n := 0 - - // for i := 0; i+contextSize <= len(tokens); i += stride { - for i := 0; i < len(tokens); i += stride { - if b.Size() == batchSize { - e.jobs <- *b - - b = newBatch(batchSize) - } - - j := min(i+contextSize, len(tokens)) - - if j-i < contextSize { - break // don't add jobs with partial windows - } - - // if j < seen { - // break // don't add jobs with no new tokens - // } - - b.AddJob(Job{ - Document: uid, - Position: n, - Tokens: tokens[i:j], - Seen: seen - i, - }) - - seen = i + contextSize - - n++ - } - - if b.Size() > 0 { - e.jobs <- *b - } - - e.scheduled.Add(int64(n)) -} - -func (e *Evaluator[R]) execute(j *batch, device int) { - if j.Size() != 1 { - panic("unimplemented") - } - - job := j.jobs[0] - - if job.Seen < 1 { - panic("empty context") - } - - m := e.models[device] - - logits := make([][]float32, 0, len(job.Tokens)) - - if _, err := m.Generate(job.Tokens, 0, &logits); err != nil { - panic(err) // TODO handle - } - - l := logits[job.Seen-1 : len(logits)-1] - t := toInt(job.Tokens[job.Seen:]) - - r := e.callback(job, l, t) - - e.results <- r - - return -} diff --git a/llm/perplexity_test.go b/llm/perplexity_test.go deleted file mode 100644 index fa974a7..0000000 --- a/llm/perplexity_test.go +++ /dev/null @@ -1,45 +0,0 @@ -package llm - -import ( - "fmt" - "testing" -) - -func TestEvaluator_estimateJobs(t *testing.T) { - type gold struct { - tokens int - window int - stride int - expected int - } - - tests := []gold{ - {tokens: 20, window: 10, stride: 3, expected: 4}, - {tokens: 20, window: 10, stride: 4, expected: 3}, - {tokens: 20, window: 10, stride: 5, expected: 3}, - - {tokens: 0, window: 1, stride: 1, expected: 0}, - {tokens: 1, window: 1, stride: 1, expected: 1}, - - {tokens: 0, window: 1024, stride: 1, expected: 0}, - {tokens: 1, window: 1, stride: 1024, expected: 1}, - } - - e := NewEvaluator[any](nil, nil, nil) - - for _, tt := range tests { - t.Run( - fmt.Sprintf("tokens%d_window%d_stride%d", tt.tokens, tt.window, tt.stride), - - func(t *testing.T) { - tokens := make([]int64, tt.tokens) - - got := e.estimateJobs(tokens, tt.window, tt.stride) - - if got != tt.expected { - t.Errorf("expected %d but got (%d, %d, %d) = %d", tt.expected, tt.tokens, tt.window, tt.stride, got) - } - }, - ) - } -} diff --git a/llm/run.go b/llm/run.go new file mode 100644 index 0000000..df48ff6 --- /dev/null +++ b/llm/run.go @@ -0,0 +1,159 @@ +package llm + +import ( + "fmt" + "sync" + "time" + + "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/tui" +) + +func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int) error { + devices := make([]int, len(e.models)) + + for i := range len(devices) { + devices[i] = i + } + + devicePool := newPool(devices...) + + var wg sync.WaitGroup + + for range e.numWorkers { + wg.Add(1) + + go func() { + defer wg.Done() + + for b := range e.jobs { + device := devicePool.Acquire() + + func() { + defer devicePool.Release(device) + + defer func() { + if r := recover(); r != nil { + fmt.Println("HOUSTON") // TODO handle + } + }() + + e.execute(&b, device) + }() + + e.completed.Add(int64(b.Size())) + } + }() + } + + pb := tui.NewProgressBar(title, 20, 0, time.Now()) + + pb.Start(1*time.Second, func() int { + return int(e.completed.Load()) + }) + + defer pb.Close() + + n := 0 + + for d := range data.Texts("text") { + tokens := toInt64(e.tokenizer.Tokenize(d)) + + e.schedule(n, tokens, window, stride, e.batchSize) + + pb.SetTotal(pb.Total() + e.estimateJobs(tokens, window, stride)) + + n++ + } + + close(e.jobs) + + wg.Wait() + + close(e.results) + + return nil +} + +func (e *Evaluator[R]) estimateJobs(tokens []int64, window, stride int) int { + if len(tokens) < window { + return 0 + } + + windows := ((len(tokens) - window) / stride) + 1 + // jobs := (windows + batchSize - 1) / e.batchSize + + return windows +} + +func (e *Evaluator[R]) schedule(uid int, tokens []int64, contextSize, stride, batchSize int) { + b := newBatch(batchSize) + + seen := 1 // first token as context + n := 0 + + // for i := 0; i+contextSize <= len(tokens); i += stride { + for i := 0; i < len(tokens); i += stride { + if b.Size() == batchSize { + e.jobs <- *b + + b = newBatch(batchSize) + } + + j := min(i+contextSize, len(tokens)) + + if j-i < contextSize { + break // don't add jobs with partial windows + } + + // if j < seen { + // break // don't add jobs with no new tokens + // } + + b.AddJob(Job{ + Document: uid, + Position: n, + Tokens: tokens[i:j], + Seen: seen - i, + }) + + seen = i + contextSize + + n++ + } + + if b.Size() > 0 { + e.jobs <- *b + } + + e.scheduled.Add(int64(n)) +} + +func (e *Evaluator[R]) execute(j *batch, device int) { + if j.Size() != 1 { + panic("unimplemented") + } + + job := j.jobs[0] + + if job.Seen < 1 { + panic("empty context") + } + + m := e.models[device] + + logits := make([][]float32, 0, len(job.Tokens)) + + if _, err := m.Generate(job.Tokens, 0, &logits); err != nil { + panic(err) // TODO handle + } + + l := logits[job.Seen-1 : len(logits)-1] + t := toInt(job.Tokens[job.Seen:]) + + r := e.callback(job, l, t) + + e.results <- r + + return +} diff --git a/llm/run_test.go b/llm/run_test.go new file mode 100644 index 0000000..fa974a7 --- /dev/null +++ b/llm/run_test.go @@ -0,0 +1,45 @@ +package llm + +import ( + "fmt" + "testing" +) + +func TestEvaluator_estimateJobs(t *testing.T) { + type gold struct { + tokens int + window int + stride int + expected int + } + + tests := []gold{ + {tokens: 20, window: 10, stride: 3, expected: 4}, + {tokens: 20, window: 10, stride: 4, expected: 3}, + {tokens: 20, window: 10, stride: 5, expected: 3}, + + {tokens: 0, window: 1, stride: 1, expected: 0}, + {tokens: 1, window: 1, stride: 1, expected: 1}, + + {tokens: 0, window: 1024, stride: 1, expected: 0}, + {tokens: 1, window: 1, stride: 1024, expected: 1}, + } + + e := NewEvaluator[any](nil, nil, nil) + + for _, tt := range tests { + t.Run( + fmt.Sprintf("tokens%d_window%d_stride%d", tt.tokens, tt.window, tt.stride), + + func(t *testing.T) { + tokens := make([]int64, tt.tokens) + + got := e.estimateJobs(tokens, tt.window, tt.stride) + + if got != tt.expected { + t.Errorf("expected %d but got (%d, %d, %d) = %d", tt.expected, tt.tokens, tt.window, tt.stride, got) + } + }, + ) + } +} -- cgit v1.3.1