summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-03-17 20:16:39 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-03-17 20:16:39 +0100
commit57fb765bdc5819f8a367ad75c731a50da68b8d85 (patch)
tree730647fc08847624c138a83f3e8329b6f4755ea6 /llm
parent035eef3308edd4ec581061440337aaa10a37f573 (diff)
Update perplexity evaluator
* Introduce results channel * Introduce jobs channel * Generalize evaluation via callback
Diffstat (limited to 'llm')
-rw-r--r--llm/cmd/eval/main.go (renamed from llm/cmd/ppl/ppl.go)19
-rw-r--r--llm/cmd/eval/ppl.go80
-rw-r--r--llm/evaluator.go43
-rw-r--r--llm/job.go37
-rw-r--r--llm/perplexity.go247
-rw-r--r--llm/perplexity_test.go45
-rw-r--r--llm/pool.go1
-rw-r--r--llm/result.go6
-rw-r--r--llm/stat.go34
9 files changed, 304 insertions, 208 deletions
diff --git a/llm/cmd/ppl/ppl.go b/llm/cmd/eval/main.go
index 6912e48..05b3f96 100644
--- a/llm/cmd/ppl/ppl.go
+++ b/llm/cmd/eval/main.go
@@ -1,33 +1,16 @@
package main
import (
- "fmt"
"log"
"github.com/jonasknobloch/mbpe"
"go.jknobloc.com/x/dataset"
"go.jknobloc.com/x/gpt2"
- "go.jknobloc.com/x/llm"
)
func main() {
- d := data()
- m := model()
- t := tokenizer()
-
- e := llm.NewEvaluator()
-
- e.SetTokenizer(t)
- e.AddModel(m)
-
- ppl, err := e.Perplexity(d, 1024, 512)
-
- if err != nil {
- log.Fatal(err)
- }
-
- fmt.Println(ppl)
+ perplexity()
}
func data() *dataset.ParquetReader {
diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go
new file mode 100644
index 0000000..fdbcff3
--- /dev/null
+++ b/llm/cmd/eval/ppl.go
@@ -0,0 +1,80 @@
+package main
+
+import (
+ "fmt"
+ "log"
+ "math"
+ "strings"
+ "sync"
+
+ "go.jknobloc.com/x/dataset"
+ "go.jknobloc.com/x/llm"
+)
+
+type pplResult struct {
+ v float64
+ n int
+}
+
+func perplexity() {
+ d := data()
+ m := model()
+ t := tokenizer()
+
+ e := llm.NewEvaluator[pplResult]()
+
+ e.SetTokenizer(t)
+ e.AddModel(m)
+
+ e.SetCallback(func(job llm.Job, logits [][]float32, tokens []int) pplResult {
+ p, n := llm.NegLogLikelihood(logits, tokens)
+
+ return pplResult{
+ v: p,
+ n: n,
+ }
+ })
+
+ results := e.Results()
+
+ var wg sync.WaitGroup
+
+ total := float64(0)
+ n := 0
+
+ wg.Add(1)
+
+ go func() {
+ defer wg.Done()
+
+ for r := range results {
+ total += r.v
+ n += r.n
+ }
+ }()
+
+ if err := e.Run(d, 1024, 512); err != nil {
+ log.Fatal(err)
+ }
+
+ wg.Wait()
+
+ avg := total / float64(n)
+ ppl := math.Exp(avg)
+
+ fmt.Println(ppl)
+}
+
+func joined() dataset.Reader {
+ miniPile := data()
+
+ docs := make([]string, 0)
+
+ for d := range miniPile.Texts("text") {
+ docs = append(docs, d)
+ }
+
+ j := dataset.NewStringReader(strings.Join(docs, "\n\n"))
+
+ return j
+}
diff --git a/llm/evaluator.go b/llm/evaluator.go
index 9faf483..24fe2d4 100644
--- a/llm/evaluator.go
+++ b/llm/evaluator.go
@@ -1,26 +1,47 @@
package llm
-import "sync"
+import (
+ "sync/atomic"
+)
-type Evaluator struct {
- mutex sync.RWMutex
+type Evaluator[R any] struct {
models []Causal
tokenizer Tokenizer
- results []result
- jobs int
+
+ batchSize int
+ numWorkers int
+
+ scheduled atomic.Int64
+ completed atomic.Int64
+
+ jobs chan batch
+ results chan R
+
+ callback func(job Job, logits [][]float32, tokens []int) R
}
-func NewEvaluator() *Evaluator {
- return &Evaluator{
- models: make([]Causal, 0),
- results: make([]result, 0),
+func NewEvaluator[R any]() *Evaluator[R] {
+ return &Evaluator[R]{
+ models: make([]Causal, 0),
+ batchSize: 1, // TODO arg
+ numWorkers: 4, // TODO arg
+ jobs: make(chan batch, 1024),
+ results: make(chan R, 1024),
}
}
-func (e *Evaluator) AddModel(model Causal) {
+func (e *Evaluator[R]) AddModel(model Causal) {
e.models = append(e.models, model)
}
-func (e *Evaluator) SetTokenizer(tokenizer Tokenizer) {
+func (e *Evaluator[R]) SetTokenizer(tokenizer Tokenizer) {
e.tokenizer = tokenizer
}
+
+func (e *Evaluator[R]) SetCallback(callback func(job Job, logits [][]float32, tokens []int) R) {
+ e.callback = callback
+}
+
+func (e *Evaluator[R]) Results() chan R {
+ return e.results
+}
diff --git a/llm/job.go b/llm/job.go
index 204d324..d922190 100644
--- a/llm/job.go
+++ b/llm/job.go
@@ -1,19 +1,30 @@
package llm
-type job struct {
- positions []int
- tokens [][]int64
- seen []int
- results []result
- debug [][2]int
+type Job struct {
+ Document int
+ Position int
+ Tokens []int64
+ Seen int
}
-func newJob(batchSize int) *job {
- return &job{
- positions: make([]int, 0, batchSize),
- tokens: make([][]int64, 0, batchSize),
- seen: make([]int, 0, batchSize),
- results: make([]result, 0, batchSize),
- debug: make([][2]int, 0),
+type batch struct {
+ jobs []Job
+}
+
+func newBatch(capacity int) *batch {
+ return &batch{
+ jobs: make([]Job, 0, capacity),
+ }
+}
+
+func (b *batch) Size() int {
+ return len(b.jobs)
+}
+
+func (b *batch) AddJob(job Job) {
+ if b.Size() == cap(b.jobs) {
+ panic("batch full")
}
+
+ b.jobs = append(b.jobs, job)
}
diff --git a/llm/perplexity.go b/llm/perplexity.go
index 2240f2d..ab31889 100644
--- a/llm/perplexity.go
+++ b/llm/perplexity.go
@@ -2,234 +2,161 @@ package llm
import (
"context"
- "log"
- "math"
- "slices"
- "strings"
+ "fmt"
"sync"
"time"
- "github.com/jonasknobloch/mbpe"
-
"go.jknobloc.com/x/dataset"
+ "go.jknobloc.com/x/tui"
)
-func (e *Evaluator) Perplexity(data *dataset.ParquetReader, window, stride int) (float64, error) {
- docs := make([]string, 0)
-
- n := 0
-
- for d := range data.Texts("text") {
- if n > 5 {
- break
- }
-
- docs = append(docs, d)
-
- n++
- }
-
- tokens := toInt64(e.tokenizer.Tokenize(strings.Join(docs, "\n\n")))[:10240] // TODO performance
-
- batchSize := 1
+func (e *Evaluator[R]) Run(data dataset.Reader, window, stride int) error {
+ devices := make([]int, len(e.models))
- if len(tokens) < window {
- return 0, nil // TODO handle
+ for i := range len(devices) {
+ devices[i] = i
}
- windows := ((len(tokens) - window) / stride) + 1
- jobs := (windows + batchSize - 1) / batchSize
+ devicePool := newPool(devices...)
- pb := mbpe.NewProgressBar("Perplexity", 20, jobs, time.Now())
-
- ctx, cancel := context.WithCancel(context.Background())
-
- defer cancel()
+ var wg sync.WaitGroup
- done := make(chan struct{})
+ for range e.numWorkers {
+ wg.Add(1)
- go func(ctx context.Context) {
- main:
- for {
- select {
- case <-ctx.Done():
- break main
- default:
- time.Sleep(time.Second * 1)
+ go func() {
+ defer wg.Done()
- e.mutex.RLock()
+ for b := range e.jobs {
+ device := devicePool.Acquire()
- j := e.jobs
+ func() {
+ defer devicePool.Release(device)
- e.mutex.RUnlock()
+ defer func() {
+ if r := recover(); r != nil {
+ fmt.Println("HOUSTON") // TODO handle
+ }
+ }()
- pb.Update(j)
- pb.Print()
+ e.execute(&b, device)
+ }()
- if j >= jobs {
- break main
- }
+ e.completed.Add(int64(b.Size()))
}
- }
-
- pb.Finish()
-
- close(done)
- }(ctx)
-
- if err := e.schedule(tokens, window, stride, batchSize); err != nil {
- log.Fatal(err)
+ }()
}
- <-done
+ ctx, cancel := context.WithCancel(context.Background())
- totalNLL := float64(0)
- totalTokens := 0
+ defer cancel()
- for _, r := range e.results {
- totalNLL += r.nll
- totalTokens += r.n
- }
+ pb := tui.NewProgressBar("Perplexity", 20, 0, time.Now())
- average := totalNLL / float64(totalTokens)
+ go pb.Watch(ctx, 1*time.Second, func() int {
+ return int(e.completed.Load())
+ })
- return math.Exp(average), nil
-}
+ n := 0
-func (e *Evaluator) schedule(tokens []int64, contextSize, stride, batchSize int) error {
- jobs := make(chan *job)
+ for d := range data.Texts("text") {
+ tokens := toInt64(e.tokenizer.Tokenize(d))
- var wg sync.WaitGroup
+ e.schedule(n, tokens, window, stride, e.batchSize)
- devices := make([]int, len(e.models))
+ pb.SetTotal(pb.Total() + e.estimateJobs(tokens, window, stride))
- for i := range len(devices) {
- devices[i] = i
+ n++
}
- devicePool := newPool[int](devices...)
+ close(e.jobs)
- for d := 0; d < devicePool.Len(); d++ {
- wg.Add(1)
-
- go func() {
- defer wg.Done()
-
- for j := range jobs {
- device := devicePool.Acquire()
-
- e.execute(j, device)
-
- devicePool.Release(device)
-
- e.mutex.Lock()
+ wg.Wait()
- for _, p := range j.results {
- e.results = append(e.results, p)
- }
+ close(e.results)
- e.jobs++
+ return nil
+}
- e.mutex.Unlock()
- }
- }()
+func (e *Evaluator[R]) estimateJobs(tokens []int64, window, stride int) int {
+ if len(tokens) < window {
+ return 0
}
- j := newJob(batchSize)
+ windows := ((len(tokens) - window) / stride) + 1
+ // jobs := (windows + batchSize - 1) / e.batchSize
- // 0 to 1023: full logits
- // 512 to 1535: 1024 upwards
- // 1024 to 2047: 1536 upwards
- // ...
+ 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 < len(tokens); i += stride {
- for i := 0; i+contextSize <= len(tokens); i += stride {
- // if (len(j.positions) == batchSize) || i+stride > len(tokens) {
- if len(j.positions) == batchSize {
- jobs <- j
+ // for i := 0; i+contextSize <= len(tokens); i += stride {
+ for i := 0; i < len(tokens); i += stride {
+ if b.Size() == batchSize {
+ e.jobs <- *b
- j = newJob(batchSize)
+ b = newBatch(batchSize)
}
- j.positions = append(j.positions, n)
- j.tokens = append(j.tokens, tokens[i:min(i+contextSize, len(tokens))]) // TODO verify
- j.seen = append(j.seen, seen-i)
+ 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 len(j.positions) > 0 {
- jobs <- j
+ if b.Size() > 0 {
+ e.jobs <- *b
}
- close(jobs)
-
- wg.Wait()
-
- return nil
+ e.scheduled.Add(int64(n))
}
-func (e *Evaluator) execute(j *job, device int) {
- if len(j.positions) != 1 {
+func (e *Evaluator[R]) execute(j *batch, device int) {
+ if j.Size() != 1 {
panic("unimplemented")
}
- if j.seen[0] < 1 {
+ job := j.jobs[0]
+
+ if job.Seen < 1 {
panic("empty context")
}
m := e.models[device]
- logits := make([][]float32, 0, len(j.tokens[0]))
-
- // fmt.Println("executing job", j.positions[0])
+ logits := make([][]float32, 0, len(job.Tokens))
- if _, err := m.Generate(j.tokens[0], 0, &logits); err != nil {
+ if _, err := m.Generate(job.Tokens, 0, &logits); err != nil {
panic(err) // TODO handle
}
- nllLogits := logits[j.seen[0]-1 : len(logits)-1]
- nnlTargets := toInt(j.tokens[0][j.seen[0]:])
+ l := logits[job.Seen-1 : len(logits)-1]
+ t := toInt(job.Tokens[job.Seen:])
- nll := negLogLikelihood(nllLogits, nnlTargets)
+ r := e.callback(job, l, t)
- j.results = append(j.results, result{
- nll: nll,
- n: len(nnlTargets),
- })
+ e.results <- r
return
}
-
-func negLogLikelihood(logits [][]float32, targets []int) float64 {
- if len(logits) != len(targets) {
- panic("mismatched input lengths")
- }
-
- total := float64(0)
-
- for i, target := range targets {
- maxLogit := float64(slices.Max(logits[i]))
-
- sumExp := float64(0)
-
- for _, v := range logits[i] {
- sumExp += math.Exp(float64(v) - maxLogit)
- }
-
- logSumExp := maxLogit + math.Log(sumExp)
-
- targetLogit := float64(logits[i][target])
-
- logProb := targetLogit - logSumExp
-
- total -= logProb
- }
-
- return total
-}
diff --git a/llm/perplexity_test.go b/llm/perplexity_test.go
new file mode 100644
index 0000000..81f10af
--- /dev/null
+++ b/llm/perplexity_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]()
+
+ 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/pool.go b/llm/pool.go
index 294b64f..68b517d 100644
--- a/llm/pool.go
+++ b/llm/pool.go
@@ -28,6 +28,7 @@ func (p *pool[K]) Len() int {
func (p *pool[K]) Acquire() K {
p.mutex.Lock()
+
defer p.mutex.Unlock()
for {
diff --git a/llm/result.go b/llm/result.go
deleted file mode 100644
index c77f68d..0000000
--- a/llm/result.go
+++ /dev/null
@@ -1,6 +0,0 @@
-package llm
-
-type result struct {
- nll float64
- n int
-}
diff --git a/llm/stat.go b/llm/stat.go
new file mode 100644
index 0000000..228e9f2
--- /dev/null
+++ b/llm/stat.go
@@ -0,0 +1,34 @@
+package llm
+
+import (
+ "math"
+ "slices"
+)
+
+func NegLogLikelihood(logits [][]float32, targets []int) (float64, int) {
+ if len(logits) != len(targets) {
+ panic("mismatched input lengths")
+ }
+
+ total := float64(0)
+
+ for i, target := range targets {
+ maxLogit := float64(slices.Max(logits[i]))
+
+ sumExp := float64(0)
+
+ for _, v := range logits[i] {
+ sumExp += math.Exp(float64(v) - maxLogit)
+ }
+
+ logSumExp := maxLogit + math.Log(sumExp)
+
+ targetLogit := float64(logits[i][target])
+
+ logProb := targetLogit - logSumExp
+
+ total -= logProb
+ }
+
+ return total, len(targets)
+}